1use std::sync::Arc;
11
12use crate::protocol::{BackendMessage, FrontendMessage, TransactionStatus};
13
14use crate::connection::{Connection, ConnectionState};
15use crate::error::{PgError, PgServerError, Result};
16use crate::query::params::encode_params_text;
17use crate::query::result::CommandTag;
18use crate::query::row::{FieldDescription, Row};
19use crate::query::{read_data_row, read_row_description};
20use crate::transport::AsyncTransport;
21
22#[non_exhaustive]
32pub struct Cursor<'a> {
33 conn: &'a mut Connection,
34 portal_name: String,
35 columns: Arc<Vec<FieldDescription>>,
36 fetch_size: i32,
37 done: bool,
38 owns_transaction: bool,
40}
41
42impl<'a> Cursor<'a> {
43 #[must_use = "cursor errors should be checked"]
47 pub async fn fetch_next(&mut self) -> Result<Vec<Row>> {
48 if self.done {
49 return Ok(Vec::new());
50 }
51
52 self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
53
54 self.conn
56 .codec
57 .encode_and_write(
58 &mut self.conn.transport,
59 &FrontendMessage::Execute {
60 portal: self.portal_name.clone(),
61 max_rows: self.fetch_size,
62 },
63 )
64 .await?;
65
66 self.conn
67 .codec
68 .encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
69 .await?;
70
71 self.conn
73 .transport
74 .flush()
75 .await
76 .map_err(PgError::Transport)?;
77
78 let mut rows = Vec::new();
79
80 loop {
81 let msg = self
82 .conn
83 .codec
84 .read_message(&mut self.conn.transport)
85 .await?;
86 if self.conn.handle_async_message(&msg) {
87 continue;
88 }
89 match msg {
90 BackendMessage::RowDescription(body) => {
91 self.columns = Arc::new(read_row_description(body)?);
92 }
93 BackendMessage::DataRow(body) => {
94 let values = read_data_row(body)?;
95 rows.push(Row::new(self.columns.clone(), values));
96 }
97 BackendMessage::CommandComplete(_body) => {
98 self.done = true;
99 }
100 BackendMessage::PortalSuspended => {
101 }
103 BackendMessage::ReadyForQuery(body) => {
104 self.conn.transaction_status = TransactionStatus::from_u8(body.status())
105 .unwrap_or(TransactionStatus::Idle);
106 self.conn.state = ConnectionState::Idle;
107 break;
108 }
109 BackendMessage::ErrorResponse(body) => {
110 let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
111 self.conn.read_until_ready().await?;
112 self.conn.state = ConnectionState::Idle;
113 return Err(PgError::Server(Box::new(server_err)));
114 }
115 _ => {}
116 }
117 }
118
119 Ok(rows)
120 }
121
122 #[must_use = "cursor close errors should be checked"]
128 pub async fn close(mut self) -> Result<()> {
129 self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
130
131 self.conn
132 .codec
133 .encode_and_write(
134 &mut self.conn.transport,
135 &FrontendMessage::Close {
136 variant: b'P',
137 name: self.portal_name.clone(),
138 },
139 )
140 .await?;
141
142 self.conn
143 .codec
144 .encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
145 .await?;
146
147 self.conn
149 .transport
150 .flush()
151 .await
152 .map_err(PgError::Transport)?;
153
154 loop {
155 let msg = self
156 .conn
157 .codec
158 .read_message(&mut self.conn.transport)
159 .await?;
160 if self.conn.handle_async_message(&msg) {
161 continue;
162 }
163 match msg {
164 BackendMessage::CloseComplete => {}
165 BackendMessage::ReadyForQuery(body) => {
166 self.conn.transaction_status = TransactionStatus::from_u8(body.status())
167 .unwrap_or(TransactionStatus::Idle);
168 self.conn.state = ConnectionState::Idle;
169 break;
170 }
171 BackendMessage::ErrorResponse(body) => {
172 let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
173 self.conn.read_until_ready().await?;
174 self.conn.state = ConnectionState::Idle;
175 return Err(PgError::Server(Box::new(server_err)));
176 }
177 _ => {}
178 }
179 }
180
181 if self.owns_transaction {
183 self.conn.execute("COMMIT").await?;
184 }
185
186 self.done = true;
187 Ok(())
188 }
189
190 pub fn is_done(&self) -> bool {
192 self.done
193 }
194}
195
196#[derive(Debug)]
202enum CursorStreamState {
203 Active,
205 Done { command_tag: CommandTag },
207 Error,
209}
210
211#[non_exhaustive]
226pub struct CursorStream<'a> {
227 conn: &'a mut Connection,
228 portal_name: String,
229 columns: Arc<Vec<FieldDescription>>,
230 fetch_size: i32,
231 state: CursorStreamState,
232 buffered_rows: Vec<Row>,
234 owns_transaction: bool,
236}
237
238impl<'a> CursorStream<'a> {
239 pub(crate) fn new(
241 conn: &'a mut Connection,
242 portal_name: String,
243 columns: Arc<Vec<FieldDescription>>,
244 fetch_size: i32,
245 owns_transaction: bool,
246 ) -> Self {
247 CursorStream {
248 conn,
249 portal_name,
250 columns,
251 fetch_size,
252 state: CursorStreamState::Active,
253 buffered_rows: Vec::new(),
254 owns_transaction,
255 }
256 }
257
258 #[must_use = "cursor stream errors should be checked"]
264 pub async fn next(&mut self) -> Result<Option<Row>> {
265 loop {
266 match self.state {
267 CursorStreamState::Done { .. } | CursorStreamState::Error => {
268 if let Some(row) = self.buffered_rows.pop() {
270 return Ok(Some(row));
271 }
272 return Ok(None);
273 }
274
275 CursorStreamState::Active => {
276 if let Some(row) = self.buffered_rows.pop() {
278 return Ok(Some(row));
279 }
280
281 self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
283
284 self.conn
285 .codec
286 .encode_and_write(
287 &mut self.conn.transport,
288 &FrontendMessage::Execute {
289 portal: self.portal_name.clone(),
290 max_rows: self.fetch_size,
291 },
292 )
293 .await?;
294
295 self.conn
296 .codec
297 .encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
298 .await?;
299
300 self.conn
301 .transport
302 .flush()
303 .await
304 .map_err(PgError::Transport)?;
305
306 let mut command_tag: Option<CommandTag> = None;
308 let mut rows: Vec<Row> = Vec::new();
309
310 loop {
311 let msg = self
312 .conn
313 .codec
314 .read_message(&mut self.conn.transport)
315 .await?;
316 if self.conn.handle_async_message(&msg) {
317 continue;
318 }
319 match msg {
320 BackendMessage::RowDescription(body) => {
321 self.columns = Arc::new(read_row_description(body)?);
322 }
323 BackendMessage::DataRow(body) => {
324 let values = read_data_row(body)?;
325 rows.push(Row::new(self.columns.clone(), values));
326 }
327 BackendMessage::CommandComplete(body) => {
328 command_tag =
329 Some(CommandTag::new(body.tag().unwrap_or("").into()));
330 }
331 BackendMessage::PortalSuspended => {
332 }
334 BackendMessage::ReadyForQuery(body) => {
335 self.conn.transaction_status =
336 TransactionStatus::from_u8(body.status())
337 .unwrap_or(TransactionStatus::Idle);
338 self.conn.state = ConnectionState::Idle;
339 break;
340 }
341 BackendMessage::ErrorResponse(body) => {
342 let server_err =
343 PgServerError::from_error_body(&body).map_err(PgError::Io)?;
344 self.conn.read_until_ready().await?;
345 self.conn.state = ConnectionState::Idle;
346 self.state = CursorStreamState::Error;
347 return Err(PgError::Server(Box::new(server_err)));
348 }
349 _ => {}
350 }
351 }
352
353 if let Some(tag) = command_tag {
355 self.state = CursorStreamState::Done { command_tag: tag };
356 }
357
358 if rows.is_empty() && !self.is_done() {
362 continue;
363 }
364
365 rows.reverse();
367 self.buffered_rows = rows;
368
369 if let Some(row) = self.buffered_rows.pop() {
371 return Ok(Some(row));
372 }
373
374 return Ok(None);
376 }
377 }
378 }
379 }
380
381 pub fn columns(&self) -> &[FieldDescription] {
383 &self.columns
384 }
385
386 pub fn is_done(&self) -> bool {
388 matches!(
389 self.state,
390 CursorStreamState::Done { .. } | CursorStreamState::Error
391 )
392 }
393
394 pub fn command_tag(&self) -> Option<&CommandTag> {
396 match &self.state {
397 CursorStreamState::Done { command_tag } => Some(command_tag),
398 _ => None,
399 }
400 }
401
402 #[must_use = "consume errors should be checked"]
405 pub async fn consume(mut self) -> Result<CommandTag> {
406 while self.next().await?.is_some() {}
407 self.close_portal().await?;
408 match &self.state {
409 CursorStreamState::Done { command_tag } => Ok(command_tag.clone()),
410 _ => Ok(CommandTag::default()),
411 }
412 }
413
414 async fn close_portal(&mut self) -> Result<()> {
416 if matches!(self.state, CursorStreamState::Done { .. }) {
417 if self.owns_transaction {
419 self.conn.execute("COMMIT").await?;
420 }
421 return Ok(());
422 }
423
424 self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
425
426 self.conn
427 .codec
428 .encode_and_write(
429 &mut self.conn.transport,
430 &FrontendMessage::Close {
431 variant: b'P',
432 name: self.portal_name.clone(),
433 },
434 )
435 .await?;
436
437 self.conn
438 .codec
439 .encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
440 .await?;
441
442 self.conn
443 .transport
444 .flush()
445 .await
446 .map_err(PgError::Transport)?;
447
448 loop {
449 let msg = self
450 .conn
451 .codec
452 .read_message(&mut self.conn.transport)
453 .await?;
454 if self.conn.handle_async_message(&msg) {
455 continue;
456 }
457 match msg {
458 BackendMessage::CloseComplete => {}
459 BackendMessage::ReadyForQuery(body) => {
460 self.conn.transaction_status = TransactionStatus::from_u8(body.status())
461 .unwrap_or(TransactionStatus::Idle);
462 self.conn.state = ConnectionState::Idle;
463 break;
464 }
465 BackendMessage::ErrorResponse(body) => {
466 let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
467 self.conn.read_until_ready().await?;
468 self.conn.state = ConnectionState::Idle;
469 return Err(PgError::Server(Box::new(server_err)));
470 }
471 _ => {}
472 }
473 }
474
475 if self.owns_transaction {
477 self.conn.execute("COMMIT").await?;
478 }
479
480 self.state = CursorStreamState::Done {
481 command_tag: CommandTag::default(),
482 };
483 Ok(())
484 }
485}
486
487impl<'a> Drop for CursorStream<'a> {
488 fn drop(&mut self) {
489 if !self.is_done() {
490 self.conn.needs_recovery = true;
491 }
492 }
493}
494
495impl Connection {
500 #[must_use = "cursor errors should be checked"]
510 pub async fn query_cursor(
511 &mut self,
512 sql: &str,
513 params: &[&dyn crate::types::ToSql],
514 fetch_size: i32,
515 ) -> Result<Cursor<'_>> {
516 let need_transaction = self.transaction_status == crate::protocol::TransactionStatus::Idle;
519 if need_transaction {
520 self.query("BEGIN").await?;
522 }
523
524 self.transition(ConnectionState::ActiveExtendedQuery)?;
525
526 let param_values = encode_params_text(params)?;
527 let portal_name = format!("__pg_portal_{}", self.statement_counter);
528 self.statement_counter += 1;
529
530 self.codec
532 .encode_and_write(
533 &mut self.transport,
534 &FrontendMessage::Parse {
535 name: String::new(),
536 sql: sql.to_string(),
537 param_types: vec![],
538 },
539 )
540 .await?;
541
542 self.codec
544 .encode_and_write(
545 &mut self.transport,
546 &FrontendMessage::Bind {
547 portal: portal_name.clone(),
548 statement: String::new(),
549 param_formats: vec![crate::protocol::FormatCode::Text],
550 params: param_values,
551 result_formats: vec![crate::protocol::FormatCode::Binary],
552 },
553 )
554 .await?;
555
556 self.codec
559 .encode_and_write(
560 &mut self.transport,
561 &FrontendMessage::Describe {
562 variant: b'P',
563 name: portal_name.clone(),
564 },
565 )
566 .await?;
567
568 self.codec
570 .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
571 .await?;
572
573 self.transport.flush().await.map_err(PgError::Transport)?;
575
576 let mut columns: Option<Arc<Vec<FieldDescription>>> = None;
577
578 loop {
579 let msg = self.codec.read_message(&mut self.transport).await?;
580 if self.handle_async_message(&msg) {
581 continue;
582 }
583 match msg {
584 BackendMessage::ParseComplete => {}
585 BackendMessage::BindComplete => {}
586 BackendMessage::NoData => {
587 }
589 BackendMessage::RowDescription(body) => {
590 columns = Some(Arc::new(read_row_description(body)?));
591 }
592 BackendMessage::ReadyForQuery(body) => {
593 self.transaction_status = TransactionStatus::from_u8(body.status())
594 .unwrap_or(TransactionStatus::Idle);
595 self.state = ConnectionState::Idle;
596 break;
597 }
598 BackendMessage::ErrorResponse(body) => {
599 let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
600 self.read_until_ready().await?;
601 self.state = ConnectionState::Idle;
602 return Err(PgError::Server(Box::new(server_err)));
603 }
604 _ => {}
605 }
606 }
607
608 Ok(Cursor {
609 conn: self,
610 portal_name,
611 columns: columns.unwrap_or_default(),
612 fetch_size,
613 done: false,
614 owns_transaction: need_transaction,
615 })
616 }
617
618 #[must_use = "cursor stream errors should be checked"]
642 pub async fn query_cursor_stream(
643 &mut self,
644 sql: &str,
645 params: &[&dyn crate::types::ToSql],
646 fetch_size: i32,
647 ) -> Result<CursorStream<'_>> {
648 let need_transaction = self.transaction_status == crate::protocol::TransactionStatus::Idle;
651 if need_transaction {
652 self.query("BEGIN").await?;
653 }
654
655 self.transition(ConnectionState::ActiveExtendedQuery)?;
656
657 let param_values = encode_params_text(params)?;
658 let portal_name = format!("__pg_portal_{}", self.statement_counter);
659 self.statement_counter += 1;
660
661 self.codec
663 .encode_and_write(
664 &mut self.transport,
665 &FrontendMessage::Parse {
666 name: String::new(),
667 sql: sql.to_string(),
668 param_types: vec![],
669 },
670 )
671 .await?;
672
673 self.codec
675 .encode_and_write(
676 &mut self.transport,
677 &FrontendMessage::Bind {
678 portal: portal_name.clone(),
679 statement: String::new(),
680 param_formats: vec![crate::protocol::FormatCode::Text],
681 params: param_values,
682 result_formats: vec![crate::protocol::FormatCode::Binary],
683 },
684 )
685 .await?;
686
687 self.codec
690 .encode_and_write(
691 &mut self.transport,
692 &FrontendMessage::Describe {
693 variant: b'P',
694 name: portal_name.clone(),
695 },
696 )
697 .await?;
698
699 self.codec
701 .encode_and_write(&mut self.transport, &FrontendMessage::Sync)
702 .await?;
703
704 self.transport.flush().await.map_err(PgError::Transport)?;
706
707 let mut columns: Option<Arc<Vec<FieldDescription>>> = None;
708
709 loop {
710 let msg = self.codec.read_message(&mut self.transport).await?;
711 if self.handle_async_message(&msg) {
712 continue;
713 }
714 match msg {
715 BackendMessage::ParseComplete => {}
716 BackendMessage::BindComplete => {}
717 BackendMessage::NoData => {
718 }
720 BackendMessage::RowDescription(body) => {
721 columns = Some(Arc::new(read_row_description(body)?));
722 }
723 BackendMessage::ReadyForQuery(body) => {
724 self.transaction_status = TransactionStatus::from_u8(body.status())
725 .unwrap_or(TransactionStatus::Idle);
726 self.state = ConnectionState::Idle;
727 break;
728 }
729 BackendMessage::ErrorResponse(body) => {
730 let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
731 self.read_until_ready().await?;
732 self.state = ConnectionState::Idle;
733 return Err(PgError::Server(Box::new(server_err)));
734 }
735 _ => {}
736 }
737 }
738
739 Ok(CursorStream::new(
740 self,
741 portal_name,
742 columns.unwrap_or_default(),
743 fetch_size,
744 need_transaction,
745 ))
746 }
747}