1use std::collections::HashMap;
2use std::sync::Arc;
3
4use async_trait::async_trait;
5use datafusion::arrow::datatypes::DataType;
6use datafusion::common::ParamValues;
7use datafusion::logical_expr::LogicalPlan;
8use datafusion::prelude::*;
9use datafusion::sql::parser::Statement;
10use datafusion::sql::sqlparser;
11use log::info;
12use pgwire::api::auth::StartupHandler;
13use pgwire::api::auth::noop::NoopStartupHandler;
14use pgwire::api::cancel::{CancelHandler, DefaultCancelHandler};
15use pgwire::api::portal::{Format, Portal};
16use pgwire::api::query::{ExtendedQueryHandler, SimpleQueryHandler};
17use pgwire::api::results::{FieldInfo, Response, Tag};
18use pgwire::api::stmt::QueryParser;
19use pgwire::api::store::PortalStore;
20use pgwire::api::{
21 ClientInfo, ClientPortalStore, ConnectionManager, ErrorHandler, PgWireServerHandlers, Type,
22};
23use pgwire::error::{PgWireError, PgWireResult};
24use pgwire::messages::PgWireBackendMessage;
25use pgwire::types::format::FormatOptions;
26
27use crate::hooks::QueryHook;
28use crate::hooks::cursor::CursorStatementHook;
29use crate::hooks::set_show::SetShowHook;
30use crate::hooks::transactions::TransactionStatementHook;
31use crate::{client, planner};
32use arrow_pg::datatypes::df;
33use arrow_pg::datatypes::{arrow_schema_to_pg_fields, into_pg_type};
34use datafusion_pg_catalog::sql::PostgresCompatibilityParser;
35
36pub struct SimpleStartupHandler {
38 connection_manager: Arc<ConnectionManager>,
39}
40
41#[async_trait::async_trait]
42impl NoopStartupHandler for SimpleStartupHandler {
43 fn connection_manager(&self) -> Option<Arc<ConnectionManager>> {
44 Some(self.connection_manager.clone())
45 }
46}
47
48pub struct HandlerFactory {
49 pub session_service: Arc<DfSessionService>,
50 cancel_handler: Arc<DefaultCancelHandler>,
51 startup_handler: Arc<SimpleStartupHandler>,
52}
53
54impl HandlerFactory {
55 pub fn new(session_context: Arc<SessionContext>) -> Self {
56 let session_service = Arc::new(DfSessionService::new(session_context));
57 let connection_manager = Arc::new(ConnectionManager::new());
58 HandlerFactory {
59 session_service,
60 cancel_handler: Arc::new(DefaultCancelHandler::new(connection_manager.clone())),
61 startup_handler: Arc::new(SimpleStartupHandler {
62 connection_manager: connection_manager.clone(),
63 }),
64 }
65 }
66
67 pub fn new_with_hooks(
68 session_context: Arc<SessionContext>,
69 query_hooks: Vec<Arc<dyn QueryHook>>,
70 ) -> Self {
71 let session_service = Arc::new(DfSessionService::new_with_hooks(
72 session_context,
73 query_hooks,
74 ));
75 let connection_manager = Arc::new(ConnectionManager::new());
76 HandlerFactory {
77 session_service,
78 cancel_handler: Arc::new(DefaultCancelHandler::new(connection_manager.clone())),
79 startup_handler: Arc::new(SimpleStartupHandler {
80 connection_manager: connection_manager.clone(),
81 }),
82 }
83 }
84}
85
86impl PgWireServerHandlers for HandlerFactory {
87 fn simple_query_handler(&self) -> Arc<impl SimpleQueryHandler> {
88 self.session_service.clone()
89 }
90
91 fn extended_query_handler(&self) -> Arc<impl ExtendedQueryHandler> {
92 self.session_service.clone()
93 }
94
95 fn startup_handler(&self) -> Arc<impl StartupHandler> {
96 self.startup_handler.clone()
97 }
98
99 fn error_handler(&self) -> Arc<impl ErrorHandler> {
100 Arc::new(LoggingErrorHandler)
101 }
102
103 fn cancel_handler(&self) -> Arc<impl CancelHandler> {
104 self.cancel_handler.clone()
105 }
106}
107
108struct LoggingErrorHandler;
109
110impl ErrorHandler for LoggingErrorHandler {
111 fn on_error<C>(&self, _client: &C, error: &mut PgWireError)
112 where
113 C: ClientInfo,
114 {
115 info!("Sending error: {error}")
116 }
117}
118
119pub struct DfSessionService {
121 session_context: Arc<SessionContext>,
122 parser: Arc<Parser>,
123 query_hooks: Vec<Arc<dyn QueryHook>>,
124}
125
126impl DfSessionService {
127 pub fn new(session_context: Arc<SessionContext>) -> DfSessionService {
128 let hooks: Vec<Arc<dyn QueryHook>> = vec![
129 Arc::new(CursorStatementHook),
130 Arc::new(SetShowHook),
131 Arc::new(TransactionStatementHook),
132 ];
133 Self::new_with_hooks(session_context, hooks)
134 }
135
136 pub fn new_with_hooks(
137 session_context: Arc<SessionContext>,
138 query_hooks: Vec<Arc<dyn QueryHook>>,
139 ) -> DfSessionService {
140 let parser = Arc::new(Parser {
141 session_context: session_context.clone(),
142 sql_parser: PostgresCompatibilityParser::new(),
143 query_hooks: query_hooks.clone(),
144 });
145 DfSessionService {
146 session_context,
147 parser,
148 query_hooks,
149 }
150 }
151}
152
153#[async_trait]
154impl SimpleQueryHandler for DfSessionService {
155 async fn do_query<C>(&self, client: &mut C, query: &str) -> PgWireResult<Vec<Response>>
156 where
157 C: ClientInfo
158 + ClientPortalStore
159 + futures::Sink<PgWireBackendMessage>
160 + Unpin
161 + Send
162 + Sync,
163 C::PortalStore: PortalStore,
164 C::Error: std::fmt::Debug,
165 PgWireError: From<<C as futures::Sink<PgWireBackendMessage>>::Error>,
166 {
167 log::debug!("Received query: {query}");
168
169 let statements = self
170 .parser
171 .sql_parser
172 .parse(query)
173 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
174
175 if statements.is_empty() {
177 return Ok(vec![Response::EmptyQuery]);
178 }
179
180 let mut results = vec![];
181 'stmt: for statement in statements {
182 for hook in &self.query_hooks {
184 if let Some(result) = hook
185 .handle_simple_query(&statement, &self.session_context, client)
186 .await
187 {
188 results.push(result?);
189 continue 'stmt;
190 }
191 }
192
193 let df_result = {
194 let query = statement.to_string();
195
196 let timeout = client::get_statement_timeout(client);
197 if let Some(timeout_duration) = timeout {
198 tokio::time::timeout(timeout_duration, self.session_context.sql(&query))
199 .await
200 .map_err(|_| {
201 PgWireError::UserError(Box::new(pgwire::error::ErrorInfo::new(
202 "ERROR".to_string(),
203 "57014".to_string(), "canceling statement due to statement timeout".to_string(),
205 )))
206 })?
207 } else {
208 self.session_context.sql(&query).await
209 }
210 };
211
212 let df = match df_result {
214 Ok(df) => df,
215 Err(e) => {
216 return Err(PgWireError::ApiError(Box::new(e)));
217 }
218 };
219
220 if matches!(statement, sqlparser::ast::Statement::Insert(_)) {
221 let resp = map_rows_affected_for_insert(&df).await?;
222 results.push(resp);
223 } else {
224 let format_options =
226 Arc::new(FormatOptions::from_client_metadata(client.metadata()));
227 let resp =
228 df::encode_dataframe(df, &Format::UnifiedText, Some(format_options)).await?;
229 results.push(Response::Query(resp));
230 }
231 }
232 Ok(results)
233 }
234}
235
236#[async_trait]
237impl ExtendedQueryHandler for DfSessionService {
238 type Statement = (String, Option<(sqlparser::ast::Statement, LogicalPlan)>);
239 type QueryParser = Parser;
240
241 fn query_parser(&self) -> Arc<Self::QueryParser> {
242 self.parser.clone()
243 }
244
245 async fn do_query<C>(
246 &self,
247 client: &mut C,
248 portal: &Portal<Self::Statement>,
249 _max_rows: usize,
250 ) -> PgWireResult<Response>
251 where
252 C: ClientInfo
253 + ClientPortalStore
254 + futures::Sink<PgWireBackendMessage>
255 + Unpin
256 + Send
257 + Sync,
258 C::PortalStore: PortalStore,
259 C::Error: std::fmt::Debug,
260 PgWireError: From<<C as futures::Sink<PgWireBackendMessage>>::Error>,
261 {
262 let query = &portal.statement.statement.0;
263 log::debug!("Received execute extended query: {query}");
264 if !self.query_hooks.is_empty()
266 && let (_, Some((statement, plan))) = &portal.statement.statement
267 {
268 let param_types = planner::get_inferred_parameter_types(plan)
270 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
271
272 let param_values: ParamValues =
273 df::deserialize_parameters(portal, &ordered_param_types(¶m_types))?;
274
275 for hook in &self.query_hooks {
276 if let Some(result) = hook
277 .handle_extended_query(
278 statement,
279 plan,
280 ¶m_values,
281 &self.session_context,
282 client,
283 )
284 .await
285 {
286 return result;
287 }
288 }
289 }
290
291 if let (_, Some((statement, plan))) = &portal.statement.statement {
292 let param_types = planner::get_inferred_parameter_types(plan)
293 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
294
295 let param_values =
296 df::deserialize_parameters(portal, &ordered_param_types(¶m_types))?;
297
298 let plan = plan
299 .clone()
300 .replace_params_with_values(¶m_values)
301 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
302 let optimised = self
303 .session_context
304 .state()
305 .optimize(&plan)
306 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
307
308 let dataframe = {
309 let timeout = client::get_statement_timeout(client);
310 if let Some(timeout_duration) = timeout {
311 tokio::time::timeout(
312 timeout_duration,
313 self.session_context.execute_logical_plan(optimised),
314 )
315 .await
316 .map_err(|_| {
317 PgWireError::UserError(Box::new(pgwire::error::ErrorInfo::new(
318 "ERROR".to_string(),
319 "57014".to_string(), "canceling statement due to statement timeout".to_string(),
321 )))
322 })?
323 .map_err(|e| PgWireError::ApiError(Box::new(e)))?
324 } else {
325 self.session_context
326 .execute_logical_plan(optimised)
327 .await
328 .map_err(|e| PgWireError::ApiError(Box::new(e)))?
329 }
330 };
331
332 if matches!(statement, sqlparser::ast::Statement::Insert(_)) {
333 let resp = map_rows_affected_for_insert(&dataframe).await?;
334
335 Ok(resp)
336 } else {
337 let format_options =
339 Arc::new(FormatOptions::from_client_metadata(client.metadata()));
340 let resp = df::encode_dataframe(
341 dataframe,
342 &portal.result_column_format,
343 Some(format_options),
344 )
345 .await?;
346 Ok(Response::Query(resp))
347 }
348 } else {
349 Ok(Response::EmptyQuery)
350 }
351 }
352}
353
354async fn map_rows_affected_for_insert(df: &DataFrame) -> PgWireResult<Response> {
355 let result = df
358 .clone()
359 .collect()
360 .await
361 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
362
363 let rows_affected = result
365 .first()
366 .and_then(|batch| batch.column_by_name("count"))
367 .and_then(|col| {
368 col.as_any()
369 .downcast_ref::<datafusion::arrow::array::UInt64Array>()
370 })
371 .map_or(0, |array| array.value(0) as usize);
372
373 let tag = Tag::new("INSERT").with_oid(0).with_rows(rows_affected);
375 Ok(Response::Execution(tag))
376}
377
378pub struct Parser {
379 session_context: Arc<SessionContext>,
380 sql_parser: PostgresCompatibilityParser,
381 query_hooks: Vec<Arc<dyn QueryHook>>,
382}
383
384#[async_trait]
385impl QueryParser for Parser {
386 type Statement = (String, Option<(sqlparser::ast::Statement, LogicalPlan)>);
387
388 async fn parse_sql<C>(
389 &self,
390 client: &C,
391 sql: &str,
392 _types: &[Option<Type>],
393 ) -> PgWireResult<Self::Statement>
394 where
395 C: ClientInfo + Unpin + Send + Sync,
396 {
397 log::debug!("Received parse extended query: {sql}");
398 let mut statements = self
399 .sql_parser
400 .parse(sql)
401 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
402 if statements.is_empty() {
403 return Ok((sql.to_string(), None));
404 }
405
406 let statement = statements.remove(0);
407 let query = statement.to_string();
408
409 let context = &self.session_context;
410 let state = context.state();
411
412 for hook in &self.query_hooks {
413 if let Some(logical_plan) = hook
414 .handle_extended_parse_query(&statement, context, client)
415 .await
416 {
417 return Ok((query, Some((statement, logical_plan?))));
418 }
419 }
420
421 let logical_plan = state
422 .statement_to_plan(Statement::Statement(Box::new(statement.clone())))
423 .await
424 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
425 Ok((query, Some((statement, logical_plan))))
426 }
427
428 fn get_parameter_types(&self, stmt: &Self::Statement) -> PgWireResult<Vec<Type>> {
429 if let (_, Some((_, plan))) = stmt {
430 let params = planner::get_inferred_parameter_types(plan)
431 .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
432
433 let mut param_types = Vec::with_capacity(params.len());
434 for param_type in ordered_param_types(¶ms).iter() {
435 if let Some(datatype) = param_type {
436 let pgtype = into_pg_type(datatype)?;
437 param_types.push(pgtype);
438 } else {
439 param_types.push(Type::UNKNOWN);
440 }
441 }
442
443 Ok(param_types)
444 } else {
445 Ok(vec![])
446 }
447 }
448
449 fn get_result_schema(
450 &self,
451 stmt: &Self::Statement,
452 column_format: Option<&Format>,
453 ) -> PgWireResult<Vec<FieldInfo>> {
454 if let (_, Some((_, plan))) = stmt {
455 if !matches!(plan, LogicalPlan::Ddl(_) | LogicalPlan::Dml(_)) {
456 let schema = plan.schema();
457 let fields = arrow_schema_to_pg_fields(
458 schema.as_arrow(),
459 column_format.unwrap_or(&Format::UnifiedText),
460 None,
461 )?;
462
463 Ok(fields)
464 } else {
465 Ok(vec![])
466 }
467 } else {
468 Ok(vec![])
469 }
470 }
471}
472
473fn ordered_param_types(types: &HashMap<String, Option<DataType>>) -> Vec<Option<&DataType>> {
474 let mut types = types.iter().collect::<Vec<_>>();
477 types.sort_by_key(|(key, _)| {
478 key.trim_start_matches('$')
479 .parse::<u32>()
480 .unwrap_or(u32::MAX)
481 });
482 types.into_iter().map(|pt| pt.1.as_ref()).collect()
483}
484
485#[cfg(test)]
486mod tests {
487 use datafusion::prelude::SessionContext;
488
489 use super::*;
490 use crate::testing::MockClient;
491
492 use crate::hooks::HookClient;
493
494 struct TestHook;
495
496 #[async_trait]
497 impl QueryHook for TestHook {
498 async fn handle_simple_query(
499 &self,
500 statement: &sqlparser::ast::Statement,
501 _ctx: &SessionContext,
502 _client: &mut dyn HookClient,
503 ) -> Option<PgWireResult<Response>> {
504 if statement.to_string().contains("magic") {
505 Some(Ok(Response::EmptyQuery))
506 } else {
507 None
508 }
509 }
510
511 async fn handle_extended_parse_query(
512 &self,
513 _statement: &sqlparser::ast::Statement,
514 _session_context: &SessionContext,
515 _client: &(dyn ClientInfo + Send + Sync),
516 ) -> Option<PgWireResult<LogicalPlan>> {
517 None
518 }
519
520 async fn handle_extended_query(
521 &self,
522 _statement: &sqlparser::ast::Statement,
523 _logical_plan: &LogicalPlan,
524 _params: &ParamValues,
525 _session_context: &SessionContext,
526 _client: &mut dyn HookClient,
527 ) -> Option<PgWireResult<Response>> {
528 None
529 }
530 }
531
532 #[test]
533 fn test_ordered_param_types_sorts_placeholders_numerically() {
534 let params = HashMap::from([
535 ("$1".to_string(), Some(DataType::Boolean)),
536 ("$2".to_string(), Some(DataType::Int64)),
537 ("$10".to_string(), Some(DataType::Utf8)),
538 ]);
539
540 let ordered = ordered_param_types(¶ms)
541 .into_iter()
542 .map(|ty| ty.cloned())
543 .collect::<Vec<_>>();
544
545 assert_eq!(
546 ordered,
547 vec![
548 Some(DataType::Boolean),
549 Some(DataType::Int64),
550 Some(DataType::Utf8)
551 ]
552 );
553 }
554
555 #[tokio::test]
556 async fn test_query_hooks() {
557 let hook = TestHook;
558 let ctx = SessionContext::new();
559 let mut client = MockClient::new();
560
561 let parser = PostgresCompatibilityParser::new();
563 let statements = parser.parse("SELECT magic").unwrap();
564 let stmt = &statements[0];
565
566 let result = hook.handle_simple_query(stmt, &ctx, &mut client).await;
568 assert!(result.is_some());
569
570 let statements = parser.parse("SELECT 1").unwrap();
572 let stmt = &statements[0];
573
574 let result = hook.handle_simple_query(stmt, &ctx, &mut client).await;
576 assert!(result.is_none());
577 }
578
579 #[tokio::test]
580 async fn test_multiple_statements_with_hook_continue() {
581 let session_context = Arc::new(SessionContext::new());
585
586 let hooks: Vec<Arc<dyn QueryHook>> = vec![Arc::new(TestHook)];
587 let service = DfSessionService::new_with_hooks(session_context, hooks);
588
589 let mut client = MockClient::new();
590
591 let query = "SELECT magic; SELECT 1; SELECT magic; SELECT 1";
593
594 let results =
595 <DfSessionService as SimpleQueryHandler>::do_query(&service, &mut client, query)
596 .await
597 .unwrap();
598
599 assert_eq!(results.len(), 4, "Expected 4 responses");
600
601 assert!(matches!(results[0], Response::EmptyQuery));
602 assert!(matches!(results[1], Response::Query(_)));
603 assert!(matches!(results[2], Response::EmptyQuery));
604 assert!(matches!(results[3], Response::Query(_)));
605 }
606
607 #[tokio::test]
608 async fn test_set_sends_parameter_status_via_sink() {
609 use pgwire::messages::PgWireBackendMessage;
610
611 let service = crate::testing::setup_handlers();
612 let mut client = MockClient::new();
613
614 let test_cases = vec![
615 ("SET datestyle = 'ISO, MDY'", "DateStyle", "ISO, MDY"),
616 (
617 "SET intervalstyle = 'postgres'",
618 "IntervalStyle",
619 "postgres",
620 ),
621 ("SET bytea_output = 'hex'", "bytea_output", "hex"),
622 (
623 "SET application_name = 'myapp'",
624 "application_name",
625 "myapp",
626 ),
627 ("SET search_path = 'public'", "search_path", "public"),
628 ("SET extra_float_digits = '2'", "extra_float_digits", "2"),
629 (
630 "SET TIME ZONE 'America/New_York'",
631 "TimeZone",
632 "America/New_York",
633 ),
634 ];
635
636 for (sql, expected_key, expected_value) in test_cases {
637 client.sent_messages.clear();
638
639 let responses =
640 <DfSessionService as SimpleQueryHandler>::do_query(&service, &mut client, sql)
641 .await
642 .unwrap();
643
644 assert!(
645 matches!(responses[0], Response::Execution(_)),
646 "Expected SET tag for {sql}"
647 );
648
649 let ps_msgs: Vec<_> = client
650 .sent_messages()
651 .iter()
652 .filter_map(|m| match m {
653 PgWireBackendMessage::ParameterStatus(ps) => Some(ps),
654 _ => None,
655 })
656 .collect();
657
658 assert_eq!(ps_msgs.len(), 1, "Expected 1 ParameterStatus for {sql}");
659 assert_eq!(ps_msgs[0].name, expected_key, "Wrong key for {sql}");
660 assert_eq!(ps_msgs[0].value, expected_value, "Wrong value for {sql}");
661 }
662 }
663
664 #[tokio::test]
665 async fn test_set_statement_timeout_no_parameter_status() {
666 use pgwire::messages::PgWireBackendMessage;
667
668 let service = crate::testing::setup_handlers();
669 let mut client = MockClient::new();
670
671 <DfSessionService as SimpleQueryHandler>::do_query(
672 &service,
673 &mut client,
674 "SET statement_timeout TO '5000ms'",
675 )
676 .await
677 .unwrap();
678
679 let has_ps = client
680 .sent_messages()
681 .iter()
682 .any(|m| matches!(m, PgWireBackendMessage::ParameterStatus(_)));
683
684 assert!(!has_ps, "statement_timeout should not send ParameterStatus");
685 }
686
687 fn assert_execution_tag(response: &Response, expected: &str) {
688 match response {
689 Response::Execution(tag) => {
690 let cc = pgwire::messages::response::CommandComplete::from(tag.clone());
691 assert_eq!(cc.tag, expected, "Unexpected execution tag");
692 }
693 other => panic!("Expected Execution response, got: {other:?}"),
694 }
695 }
696
697 async fn assert_query_response_empty(response: &mut Response) {
698 use futures::StreamExt;
699
700 let Response::Query(qr) = response else {
701 panic!("Expected Query response, got: {response:?}");
702 };
703
704 let mut count = 0;
705 while qr.data_rows().next().await.is_some() {
706 count += 1;
707 }
708 assert_eq!(count, 0, "Expected no rows from exhausted cursor");
709 }
710
711 #[tokio::test]
712 async fn test_declare_fetch_close_cursor() {
713 let service = crate::testing::setup_handlers();
714 let mut client = MockClient::new();
715
716 let responses = <DfSessionService as SimpleQueryHandler>::do_query(
717 &service,
718 &mut client,
719 "DECLARE test_cursor CURSOR FOR SELECT 1 AS col",
720 )
721 .await
722 .unwrap();
723
724 assert_eq!(responses.len(), 1);
725 assert_execution_tag(&responses[0], "DECLARE CURSOR");
726
727 let responses = <DfSessionService as SimpleQueryHandler>::do_query(
728 &service,
729 &mut client,
730 "FETCH NEXT FROM test_cursor",
731 )
732 .await
733 .unwrap();
734
735 assert_eq!(responses.len(), 1);
736 assert!(
737 matches!(&responses[0], Response::Query(_)),
738 "Expected Query response for FETCH"
739 );
740
741 let mut responses = <DfSessionService as SimpleQueryHandler>::do_query(
742 &service,
743 &mut client,
744 "FETCH NEXT FROM test_cursor",
745 )
746 .await
747 .unwrap();
748
749 assert_eq!(responses.len(), 1);
750 assert_query_response_empty(&mut responses[0]).await;
751
752 let responses = <DfSessionService as SimpleQueryHandler>::do_query(
753 &service,
754 &mut client,
755 "CLOSE test_cursor",
756 )
757 .await
758 .unwrap();
759
760 assert_eq!(responses.len(), 1);
761 assert_execution_tag(&responses[0], "CLOSE CURSOR");
762 }
763
764 #[tokio::test]
765 async fn test_fetch_nonexistent_cursor() {
766 let service = crate::testing::setup_handlers();
767 let mut client = MockClient::new();
768
769 let result = <DfSessionService as SimpleQueryHandler>::do_query(
770 &service,
771 &mut client,
772 "FETCH NEXT FROM nonexistent",
773 )
774 .await;
775
776 assert!(result.is_err());
777 }
778
779 #[tokio::test]
780 async fn test_close_all_portals() {
781 let service = crate::testing::setup_handlers();
782 let mut client = MockClient::new();
783
784 <DfSessionService as SimpleQueryHandler>::do_query(
785 &service,
786 &mut client,
787 "DECLARE c1 CURSOR FOR SELECT 1",
788 )
789 .await
790 .unwrap();
791
792 <DfSessionService as SimpleQueryHandler>::do_query(
793 &service,
794 &mut client,
795 "DECLARE c2 CURSOR FOR SELECT 2",
796 )
797 .await
798 .unwrap();
799
800 let responses =
801 <DfSessionService as SimpleQueryHandler>::do_query(&service, &mut client, "CLOSE ALL")
802 .await
803 .unwrap();
804
805 assert!(matches!(&responses[0], Response::Execution(_)),);
806
807 let result = <DfSessionService as SimpleQueryHandler>::do_query(
808 &service,
809 &mut client,
810 "FETCH NEXT FROM c1",
811 )
812 .await;
813 assert!(result.is_err(), "c1 should be closed");
814 }
815
816 #[tokio::test]
817 async fn test_fetch_forward_n() {
818 let service = crate::testing::setup_handlers();
819 let mut client = MockClient::new();
820
821 <DfSessionService as SimpleQueryHandler>::do_query(
822 &service,
823 &mut client,
824 "CREATE TABLE nums AS SELECT 1 AS n UNION ALL SELECT 2 UNION ALL SELECT 3 UNION ALL SELECT 4 UNION ALL SELECT 5",
825 )
826 .await
827 .unwrap();
828
829 <DfSessionService as SimpleQueryHandler>::do_query(
830 &service,
831 &mut client,
832 "DECLARE mycur CURSOR FOR SELECT n FROM nums ORDER BY n",
833 )
834 .await
835 .unwrap();
836
837 let responses = <DfSessionService as SimpleQueryHandler>::do_query(
838 &service,
839 &mut client,
840 "FETCH FORWARD 3 FROM mycur",
841 )
842 .await
843 .unwrap();
844
845 assert!(
846 matches!(&responses[0], Response::Query(_)),
847 "Expected Query response for FORWARD 3"
848 );
849
850 let responses = <DfSessionService as SimpleQueryHandler>::do_query(
851 &service,
852 &mut client,
853 "FETCH FORWARD ALL FROM mycur",
854 )
855 .await
856 .unwrap();
857
858 let resp_desc = match &responses[0] {
859 Response::Query(_) => "Query".to_string(),
860 Response::Execution(tag) => {
861 let cc = pgwire::messages::response::CommandComplete::from(tag.clone());
862 format!("Execution({})", cc.tag)
863 }
864 other => format!("{:?}", other),
865 };
866 assert!(
867 matches!(&responses[0], Response::Query(_)),
868 "Expected Query response for remaining rows, got: {resp_desc}"
869 );
870
871 let mut responses = <DfSessionService as SimpleQueryHandler>::do_query(
872 &service,
873 &mut client,
874 "FETCH NEXT FROM mycur",
875 )
876 .await
877 .unwrap();
878
879 assert_query_response_empty(&mut responses[0]).await;
880 }
881
882 #[tokio::test]
883 async fn test_scroll_cursor_error() {
884 let service = crate::testing::setup_handlers();
885 let mut client = MockClient::new();
886
887 <DfSessionService as SimpleQueryHandler>::do_query(
888 &service,
889 &mut client,
890 "DECLARE mycur CURSOR FOR SELECT 1",
891 )
892 .await
893 .unwrap();
894
895 let result = <DfSessionService as SimpleQueryHandler>::do_query(
896 &service,
897 &mut client,
898 "FETCH PRIOR FROM mycur",
899 )
900 .await;
901
902 assert!(result.is_err(), "PRIOR should fail on forward-only cursor");
903 }
904
905 #[tokio::test]
906 async fn test_move_cursor() {
907 let service = crate::testing::setup_handlers();
908 let mut client = MockClient::new();
909
910 <DfSessionService as SimpleQueryHandler>::do_query(
911 &service,
912 &mut client,
913 "DECLARE mycur CURSOR FOR SELECT generate_series(1, 5) AS n",
914 )
915 .await
916 .unwrap();
917
918 let responses = <DfSessionService as SimpleQueryHandler>::do_query(
919 &service,
920 &mut client,
921 "FETCH FORWARD 3 FROM mycur",
922 )
923 .await
924 .unwrap();
925
926 assert!(matches!(&responses[0], Response::Query(_)));
927 }
928}