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 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 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 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 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 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 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 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 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 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 let timeout = client::get_statement_timeout(&client);
396 assert_eq!(timeout, Some(Duration::from_millis(5000)));
397
398 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 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 let bytea_output = client.metadata().get("bytea_output").unwrap();
430 assert_eq!(bytea_output, "hex");
431
432 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 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 let bytea_output = client.metadata().get("datestyle").unwrap();
464 assert_eq!(bytea_output, "ISO, DMY");
465
466 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 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 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}