Skip to main content

datafusion_postgres/
handlers.rs

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
36/// Simple startup handler that does no authentication
37pub 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
119/// The pgwire handler backed by a datafusion `SessionContext`
120pub 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        // empty query
176        if statements.is_empty() {
177            return Ok(vec![Response::EmptyQuery]);
178        }
179
180        let mut results = vec![];
181        'stmt: for statement in statements {
182            // Call query hooks with the parsed statement
183            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(), // query_canceled error code
204                                "canceling statement due to statement timeout".to_string(),
205                            )))
206                        })?
207                } else {
208                    self.session_context.sql(&query).await
209                }
210            };
211
212            // Handle query execution errors and transaction state
213            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                // For non-INSERT queries, return a regular Query response
225                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        // Check query hooks first
265        if !self.query_hooks.is_empty()
266            && let (_, Some((statement, plan))) = &portal.statement.statement
267        {
268            // TODO: in the case where query hooks all return None, we do the param handling again later.
269            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(&param_types))?;
274
275            for hook in &self.query_hooks {
276                if let Some(result) = hook
277                    .handle_extended_query(
278                        statement,
279                        plan,
280                        &param_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(&param_types))?;
297
298            let plan = plan
299                .clone()
300                .replace_params_with_values(&param_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(), // query_canceled error code
320                            "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                // For non-INSERT queries, return a regular Query response
338                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    // For INSERT queries, we need to execute the query to get the row count
356    // and return an Execution response with the proper tag
357    let result = df
358        .clone()
359        .collect()
360        .await
361        .map_err(|e| PgWireError::ApiError(Box::new(e)))?;
362
363    // Extract count field from the first batch
364    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    // Create INSERT tag with the affected row count
374    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(&params).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    // Datafusion stores the parameters as a map.  In our case, the keys will be
475    // `$1`, `$2` etc.  The values will be the parameter types.
476    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(&params)
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        // Parse a statement that contains "magic"
562        let parser = PostgresCompatibilityParser::new();
563        let statements = parser.parse("SELECT magic").unwrap();
564        let stmt = &statements[0];
565
566        // Hook should intercept
567        let result = hook.handle_simple_query(stmt, &ctx, &mut client).await;
568        assert!(result.is_some());
569
570        // Parse a normal statement
571        let statements = parser.parse("SELECT 1").unwrap();
572        let stmt = &statements[0];
573
574        // Hook should not intercept
575        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        // Bug #227: when a hook returned a result, the code used `break 'stmt`
582        // which would exit the entire statement loop, preventing subsequent statements
583        // from being processed.
584        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        // Mix of queries with hooks and those without
592        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}