Skip to main content

datafusion_postgres/hooks/
set_show.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use datafusion::arrow::datatypes::{DataType, Field, Schema};
5use datafusion::common::{ParamValues, ToDFSchema};
6use datafusion::error::DataFusionError;
7use datafusion::logical_expr::LogicalPlan;
8use datafusion::prelude::SessionContext;
9use datafusion::sql::sqlparser::ast::{Expr, Set, Statement};
10use log::{info, warn};
11use pgwire::api::ClientInfo;
12use pgwire::api::auth::DefaultServerParameterProvider;
13use pgwire::api::results::{DataRowEncoder, FieldFormat, FieldInfo, QueryResponse, Response, Tag};
14use pgwire::error::{PgWireError, PgWireResult};
15use pgwire::messages::PgWireBackendMessage;
16use pgwire::messages::startup::ParameterStatus;
17use pgwire::types::format::FormatOptions;
18use postgres_types::Type;
19
20use crate::QueryHook;
21use crate::client;
22use crate::hooks::HookClient;
23
24#[derive(Debug)]
25pub struct SetShowHook;
26
27#[async_trait]
28impl QueryHook for SetShowHook {
29    /// called in simple query handler to return response directly
30    async fn handle_simple_query(
31        &self,
32        statement: &Statement,
33        session_context: &SessionContext,
34        client: &mut dyn HookClient,
35    ) -> Option<PgWireResult<Response>> {
36        match statement {
37            Statement::Set { .. } => {
38                try_respond_set_statements(client, statement, session_context).await
39            }
40            Statement::ShowVariable { .. }
41            | Statement::ShowStatus { .. }
42            | Statement::ShowCatalogs { .. } => {
43                try_respond_show_statements(client, statement, session_context).await
44            }
45            _ => None,
46        }
47    }
48
49    async fn handle_extended_parse_query(
50        &self,
51        stmt: &Statement,
52        _session_context: &SessionContext,
53        _client: &(dyn ClientInfo + Send + Sync),
54    ) -> Option<PgWireResult<LogicalPlan>> {
55        match stmt {
56            Statement::Set { .. } => {
57                let show_schema = Arc::new(Schema::new(Vec::<Field>::new()));
58                let result = show_schema
59                    .to_dfschema()
60                    .map(|df_schema| {
61                        LogicalPlan::EmptyRelation(datafusion::logical_expr::EmptyRelation {
62                            produce_one_row: true,
63                            schema: Arc::new(df_schema),
64                        })
65                    })
66                    .map_err(|e| PgWireError::ApiError(Box::new(e)));
67                Some(result)
68            }
69            Statement::ShowVariable { .. }
70            | Statement::ShowStatus { .. }
71            | Statement::ShowCatalogs { .. } => {
72                let show_schema =
73                    Arc::new(Schema::new(vec![Field::new("show", DataType::Utf8, false)]));
74                let result = show_schema
75                    .to_dfschema()
76                    .map(|df_schema| {
77                        LogicalPlan::EmptyRelation(datafusion::logical_expr::EmptyRelation {
78                            produce_one_row: true,
79                            schema: Arc::new(df_schema),
80                        })
81                    })
82                    .map_err(|e| PgWireError::ApiError(Box::new(e)));
83                Some(result)
84            }
85            _ => None,
86        }
87    }
88
89    async fn handle_extended_query(
90        &self,
91        statement: &Statement,
92        _logical_plan: &LogicalPlan,
93        _params: &ParamValues,
94        session_context: &SessionContext,
95        client: &mut dyn HookClient,
96    ) -> Option<PgWireResult<Response>> {
97        match statement {
98            Statement::Set { .. } => {
99                try_respond_set_statements(client, statement, session_context).await
100            }
101            Statement::ShowVariable { .. }
102            | Statement::ShowStatus { .. }
103            | Statement::ShowCatalogs { .. } => {
104                try_respond_show_statements(client, statement, session_context).await
105            }
106            _ => None,
107        }
108    }
109}
110
111fn mock_show_response(name: &str, value: &str) -> PgWireResult<QueryResponse> {
112    let fields = vec![FieldInfo::new(
113        name.to_string(),
114        None,
115        None,
116        Type::VARCHAR,
117        FieldFormat::Text,
118    )];
119
120    let row = {
121        let mut encoder = DataRowEncoder::new(Arc::new(fields.clone()));
122        encoder.encode_field(&Some(value))?;
123        Ok(encoder.take_row())
124    };
125
126    let row_stream = futures::stream::once(async move { row });
127    Ok(QueryResponse::new(Arc::new(fields), Box::pin(row_stream)))
128}
129
130async fn try_respond_set_statements(
131    client: &mut dyn HookClient,
132    statement: &Statement,
133    session_context: &SessionContext,
134) -> Option<PgWireResult<Response>> {
135    let Statement::Set(set_statement) = statement else {
136        return None;
137    };
138
139    match &set_statement {
140        Set::SingleAssignment {
141            scope: None,
142            hivevar: false,
143            variable,
144            values,
145        } => {
146            let var = variable.to_string().to_lowercase();
147            if var == "statement_timeout" {
148                let value = values[0].to_string();
149                let timeout_str = value.trim_matches('"').trim_matches('\'');
150
151                let timeout = if timeout_str == "0" || timeout_str.is_empty() {
152                    None
153                } else {
154                    // Parse timeout value (supports ms, s, min formats)
155                    let timeout_ms = if timeout_str.ends_with("ms") {
156                        timeout_str.trim_end_matches("ms").parse::<u64>()
157                    } else if timeout_str.ends_with("s") {
158                        timeout_str
159                            .trim_end_matches("s")
160                            .parse::<u64>()
161                            .map(|s| s * 1000)
162                    } else if timeout_str.ends_with("min") {
163                        timeout_str
164                            .trim_end_matches("min")
165                            .parse::<u64>()
166                            .map(|m| m * 60 * 1000)
167                    } else {
168                        // Default to milliseconds
169                        timeout_str.parse::<u64>()
170                    };
171
172                    match timeout_ms {
173                        Ok(ms) if ms > 0 => Some(std::time::Duration::from_millis(ms)),
174                        _ => None,
175                    }
176                };
177
178                client::set_statement_timeout(client, timeout);
179                return Some(Ok(Response::Execution(Tag::new("SET"))));
180            } else if matches!(
181                var.as_str(),
182                "datestyle"
183                    | "bytea_output"
184                    | "intervalstyle"
185                    | "application_name"
186                    | "extra_float_digits"
187                    | "search_path"
188            ) && !values.is_empty()
189            {
190                // postgres configuration variables
191                let value = values[0].clone();
192                if let Expr::Value(value) = value {
193                    let val_str = value.into_string().unwrap_or_else(|| "".to_string());
194                    client.metadata_mut().insert(var.clone(), val_str);
195                    if let Some((name, value)) = parameter_status_for_var(&var, &*client)
196                        && let Err(e) = client
197                            .send_message(PgWireBackendMessage::ParameterStatus(
198                                ParameterStatus::new(name, value),
199                            ))
200                            .await
201                    {
202                        return Some(Err(e));
203                    }
204                    return Some(Ok(Response::Execution(Tag::new("SET"))));
205                }
206            }
207        }
208        Set::SetTimeZone {
209            local: false,
210            value,
211        } => {
212            let tz = value.to_string();
213            let tz = tz.trim_matches('"').trim_matches('\'');
214            client::set_timezone(client, Some(tz));
215            // execution options for timezone
216            session_context
217                .state()
218                .config_mut()
219                .options_mut()
220                .execution
221                .time_zone = Some(tz.to_string());
222            let tz_value = client::get_timezone(client).unwrap_or("UTC").to_string();
223            if let Err(e) = client
224                .send_message(PgWireBackendMessage::ParameterStatus(ParameterStatus::new(
225                    "TimeZone".to_string(),
226                    tz_value,
227                )))
228                .await
229            {
230                return Some(Err(e));
231            }
232            return Some(Ok(Response::Execution(Tag::new("SET"))));
233        }
234        _ => {}
235    }
236
237    // fallback to datafusion and ignore all errors
238    if let Err(e) = execute_set_statement(session_context, statement.clone()).await {
239        warn!(
240            "SET statement {statement} is not supported by datafusion, error {e}, statement ignored",
241        );
242    }
243
244    // Always return SET success
245    Some(Ok(Response::Execution(Tag::new("SET"))))
246}
247
248fn parameter_status_for_var(
249    var: &str,
250    client: &(impl ClientInfo + ?Sized),
251) -> Option<(String, String)> {
252    let display_name = match var {
253        "datestyle" => "DateStyle",
254        "intervalstyle" => "IntervalStyle",
255        "bytea_output" => "bytea_output",
256        "application_name" => "application_name",
257        "extra_float_digits" => "extra_float_digits",
258        "search_path" => "search_path",
259        _ => return None,
260    };
261    let value = client.metadata().get(var)?.clone();
262    Some((display_name.to_string(), value))
263}
264
265async fn execute_set_statement(
266    session_context: &SessionContext,
267    statement: Statement,
268) -> Result<(), DataFusionError> {
269    let state = session_context.state();
270    let logical_plan = state
271        .statement_to_plan(datafusion::sql::parser::Statement::Statement(Box::new(
272            statement,
273        )))
274        .await
275        .and_then(|logical_plan| state.optimize(&logical_plan))?;
276
277    session_context
278        .execute_logical_plan(logical_plan)
279        .await
280        .map(|_| ())
281}
282
283async fn try_respond_show_statements(
284    client: &dyn HookClient,
285    statement: &Statement,
286    session_context: &SessionContext,
287) -> Option<PgWireResult<Response>> {
288    // Handle SHOW CATALOGS separately since it's its own variant in sqlparser 0.62+
289    if let Statement::ShowCatalogs { .. } = statement {
290        let catalogs = session_context.catalog_names();
291        let value = catalogs.join(", ");
292        return Some(mock_show_response("Catalogs", &value).map(Response::Query));
293    }
294
295    let Statement::ShowVariable { variable } = statement else {
296        return None;
297    };
298
299    let variables = variable
300        .iter()
301        .map(|v| v.value.to_lowercase())
302        .collect::<Vec<_>>();
303    let variables_ref = variables.iter().map(|s| s.as_str()).collect::<Vec<_>>();
304
305    match variables_ref.as_slice() {
306        ["time", "zone"] => {
307            let timezone = client::get_timezone(client).unwrap_or("UTC");
308            Some(mock_show_response("TimeZone", timezone).map(Response::Query))
309        }
310        ["server_version"] => {
311            let version = format!(
312                "datafusion {} on {} {}",
313                session_context.state().version(),
314                env!("CARGO_PKG_NAME"),
315                env!("CARGO_PKG_VERSION")
316            );
317            Some(mock_show_response("server_version", &version).map(Response::Query))
318        }
319        ["transaction_isolation"] => Some(
320            mock_show_response("transaction_isolation", "read uncommitted").map(Response::Query),
321        ),
322        ["catalogs"] => {
323            let catalogs = session_context.catalog_names();
324            let value = catalogs.join(", ");
325            Some(mock_show_response("Catalogs", &value).map(Response::Query))
326        }
327        ["statement_timeout"] => {
328            let timeout = client::get_statement_timeout(client);
329            let timeout_str = match timeout {
330                Some(duration) => format!("{}ms", duration.as_millis()),
331                None => "0".to_string(),
332            };
333            Some(mock_show_response("statement_timeout", &timeout_str).map(Response::Query))
334        }
335        ["transaction", "isolation", "level"] => {
336            Some(mock_show_response("transaction_isolation", "read_committed").map(Response::Query))
337        }
338        _ => {
339            let val = client
340                .metadata()
341                .get(&variables[0])
342                .map(|v| v.to_string())
343                .or_else(|| match variables[0].as_str() {
344                    "bytea_output" => Some(FormatOptions::default().bytea_output),
345                    "datestyle" => Some(FormatOptions::default().date_style),
346                    "intervalstyle" => Some(FormatOptions::default().interval_style),
347                    "extra_float_digits" => {
348                        Some(FormatOptions::default().extra_float_digits.to_string())
349                    }
350                    "application_name" => Some(
351                        DefaultServerParameterProvider::default()
352                            .application_name
353                            .unwrap_or("".to_owned()),
354                    ),
355                    "search_path" => Some(DefaultServerParameterProvider::default().search_path),
356                    _ => None,
357                });
358            if let Some(val) = val {
359                Some(mock_show_response(&variables[0], &val).map(Response::Query))
360            } else {
361                info!("Unsupported show statement: {statement}");
362                Some(mock_show_response("unsupported_show_statement", "").map(Response::Query))
363            }
364        }
365    }
366}
367
368#[cfg(test)]
369mod tests {
370    use std::time::Duration;
371
372    use datafusion::sql::sqlparser::{dialect::PostgreSqlDialect, parser::Parser};
373
374    use super::*;
375    use crate::testing::MockClient;
376
377    #[tokio::test]
378    async fn test_statement_timeout_set_and_show() {
379        let session_context = SessionContext::new();
380        let mut client = MockClient::new();
381
382        // Test setting timeout to 5000ms
383        let statement = Parser::new(&PostgreSqlDialect {})
384            .try_with_sql("set statement_timeout to '5000ms'")
385            .unwrap()
386            .parse_statement()
387            .unwrap();
388        let set_response =
389            try_respond_set_statements(&mut client, &statement, &session_context).await;
390
391        assert!(set_response.is_some());
392        assert!(set_response.unwrap().is_ok());
393
394        // Verify the timeout was set in client metadata
395        let timeout = client::get_statement_timeout(&client);
396        assert_eq!(timeout, Some(Duration::from_millis(5000)));
397
398        // Test SHOW statement_timeout
399        let statement = Parser::new(&PostgreSqlDialect {})
400            .try_with_sql("show statement_timeout")
401            .unwrap()
402            .parse_statement()
403            .unwrap();
404        let show_response =
405            try_respond_show_statements(&client, &statement, &session_context).await;
406
407        assert!(show_response.is_some());
408        assert!(show_response.unwrap().is_ok());
409    }
410
411    #[tokio::test]
412    async fn test_bytea_output_set_and_show() {
413        let session_context = SessionContext::new();
414        let mut client = MockClient::new();
415
416        // Test setting bytea_output to hex
417        let statement = Parser::new(&PostgreSqlDialect {})
418            .try_with_sql("set bytea_output = 'hex'")
419            .unwrap()
420            .parse_statement()
421            .unwrap();
422        let set_response =
423            try_respond_set_statements(&mut client, &statement, &session_context).await;
424
425        assert!(set_response.is_some());
426        assert!(set_response.unwrap().is_ok());
427
428        // Verify the value was set in client metadata
429        let bytea_output = client.metadata().get("bytea_output").unwrap();
430        assert_eq!(bytea_output, "hex");
431
432        // Test SHOW bytea_output
433        let statement = Parser::new(&PostgreSqlDialect {})
434            .try_with_sql("show bytea_output")
435            .unwrap()
436            .parse_statement()
437            .unwrap();
438        let show_response =
439            try_respond_show_statements(&client, &statement, &session_context).await;
440
441        assert!(show_response.is_some());
442        assert!(show_response.unwrap().is_ok());
443    }
444
445    #[tokio::test]
446    async fn test_date_style_set_and_show() {
447        let session_context = SessionContext::new();
448        let mut client = MockClient::new();
449
450        // Test setting dateStyle
451        let statement = Parser::new(&PostgreSqlDialect {})
452            .try_with_sql("set dateStyle = 'ISO, DMY'")
453            .unwrap()
454            .parse_statement()
455            .unwrap();
456        let set_response =
457            try_respond_set_statements(&mut client, &statement, &session_context).await;
458
459        assert!(set_response.is_some());
460        assert!(set_response.unwrap().is_ok());
461
462        // Verify the value was set in client metadata
463        let bytea_output = client.metadata().get("datestyle").unwrap();
464        assert_eq!(bytea_output, "ISO, DMY");
465
466        // Test SHOW dateStyle
467        let statement = Parser::new(&PostgreSqlDialect {})
468            .try_with_sql("show dateStyle")
469            .unwrap()
470            .parse_statement()
471            .unwrap();
472        let show_response =
473            try_respond_show_statements(&client, &statement, &session_context).await;
474
475        assert!(show_response.is_some());
476        assert!(show_response.unwrap().is_ok());
477    }
478
479    #[tokio::test]
480    async fn test_statement_timeout_disable() {
481        let session_context = SessionContext::new();
482        let mut client = MockClient::new();
483
484        // Set timeout first
485        let statement = Parser::new(&PostgreSqlDialect {})
486            .try_with_sql("set statement_timeout to '1000ms'")
487            .unwrap()
488            .parse_statement()
489            .unwrap();
490        let resp = try_respond_set_statements(&mut client, &statement, &session_context).await;
491        assert!(resp.is_some());
492        assert!(resp.unwrap().is_ok());
493
494        // Disable timeout with 0
495        let statement = Parser::new(&PostgreSqlDialect {})
496            .try_with_sql("set statement_timeout to '0'")
497            .unwrap()
498            .parse_statement()
499            .unwrap();
500        let resp = try_respond_set_statements(&mut client, &statement, &session_context).await;
501        assert!(resp.is_some());
502        assert!(resp.unwrap().is_ok());
503
504        let timeout = client::get_statement_timeout(&client);
505        assert_eq!(timeout, None);
506    }
507
508    #[tokio::test]
509    async fn test_parameter_status_sent_for_all_set_vars() {
510        use pgwire::messages::PgWireBackendMessage;
511
512        let test_cases = vec![
513            ("set bytea_output = 'escape'", "bytea_output", "escape"),
514            (
515                "set intervalstyle = 'postgres'",
516                "IntervalStyle",
517                "postgres",
518            ),
519            (
520                "set application_name = 'myapp'",
521                "application_name",
522                "myapp",
523            ),
524            ("set search_path = 'public'", "search_path", "public"),
525            ("set extra_float_digits = '2'", "extra_float_digits", "2"),
526            ("set datestyle = 'ISO, MDY'", "DateStyle", "ISO, MDY"),
527            (
528                "set time zone 'America/New_York'",
529                "TimeZone",
530                "America/New_York",
531            ),
532        ];
533
534        for (sql, expected_key, expected_value) in test_cases {
535            let session_context = SessionContext::new();
536            let mut client = MockClient::new();
537            let statement = Parser::new(&PostgreSqlDialect {})
538                .try_with_sql(sql)
539                .unwrap()
540                .parse_statement()
541                .unwrap();
542
543            let result =
544                try_respond_set_statements(&mut client, &statement, &session_context).await;
545            assert!(result.is_some(), "Expected Some for {sql}");
546            assert!(result.unwrap().is_ok(), "Expected Ok for {sql}");
547
548            let ps_msgs: Vec<_> = client
549                .sent_messages()
550                .iter()
551                .filter_map(|m| match m {
552                    PgWireBackendMessage::ParameterStatus(ps) => Some(ps),
553                    _ => None,
554                })
555                .collect();
556
557            assert_eq!(ps_msgs.len(), 1, "Expected 1 ParameterStatus for {sql}");
558            assert_eq!(ps_msgs[0].name, expected_key, "Wrong key for {sql}");
559            assert_eq!(ps_msgs[0].value, expected_value, "Wrong value for {sql}");
560        }
561    }
562
563    #[tokio::test]
564    async fn test_no_parameter_status_for_statement_timeout() {
565        use pgwire::messages::PgWireBackendMessage;
566
567        let session_context = SessionContext::new();
568        let mut client = MockClient::new();
569
570        let statement = Parser::new(&PostgreSqlDialect {})
571            .try_with_sql("set statement_timeout to '5000ms'")
572            .unwrap()
573            .parse_statement()
574            .unwrap();
575
576        let result = try_respond_set_statements(&mut client, &statement, &session_context).await;
577        assert!(result.is_some());
578        assert!(result.unwrap().is_ok());
579
580        let has_ps = client
581            .sent_messages()
582            .iter()
583            .any(|m| matches!(m, PgWireBackendMessage::ParameterStatus(_)));
584
585        assert!(!has_ps, "statement_timeout should not send ParameterStatus");
586    }
587
588    #[tokio::test]
589    async fn test_supported_show_statements_returned_columns() {
590        let session_context = SessionContext::new();
591        let client = MockClient::new();
592
593        let tests = [
594            ("show time zone", "TimeZone"),
595            ("show server_version", "server_version"),
596            ("show transaction_isolation", "transaction_isolation"),
597            ("show catalogs", "Catalogs"),
598            ("show search_path", "search_path"),
599            ("show statement_timeout", "statement_timeout"),
600            ("show transaction isolation level", "transaction_isolation"),
601        ];
602
603        for (query, expected_response_col) in tests {
604            let statement = Parser::new(&PostgreSqlDialect {})
605                .try_with_sql(&query)
606                .unwrap()
607                .parse_statement()
608                .unwrap();
609            let show_response =
610                try_respond_show_statements(&client, &statement, &session_context).await;
611
612            dbg!(query);
613            dbg!(&show_response);
614            let Some(Ok(Response::Query(show_response))) = show_response else {
615                panic!("unexpected show response");
616            };
617
618            assert_eq!(show_response.command_tag(), "SELECT");
619
620            let row_schema = show_response.row_schema();
621            assert_eq!(row_schema.len(), 1);
622            assert_eq!(row_schema[0].name(), expected_response_col);
623        }
624    }
625}