use std::sync::Arc;
use crate::protocol::{BackendMessage, FrontendMessage, TransactionStatus};
use crate::connection::{Connection, ConnectionState};
use crate::error::{PgError, PgServerError, Result};
use crate::query::params::encode_params_text;
use crate::query::result::CommandTag;
use crate::query::row::{FieldDescription, Row};
use crate::query::{read_data_row, read_row_description};
use crate::transport::AsyncTransport;
#[non_exhaustive]
pub struct Cursor<'a> {
conn: &'a mut Connection,
portal_name: String,
columns: Arc<Vec<FieldDescription>>,
fetch_size: i32,
done: bool,
owns_transaction: bool,
}
impl<'a> Cursor<'a> {
#[must_use = "cursor errors should be checked"]
pub async fn fetch_next(&mut self) -> Result<Vec<Row>> {
if self.done {
return Ok(Vec::new());
}
self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
self.conn
.codec
.encode_and_write(
&mut self.conn.transport,
&FrontendMessage::Execute {
portal: self.portal_name.clone(),
max_rows: self.fetch_size,
},
)
.await?;
self.conn
.codec
.encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
.await?;
self.conn
.transport
.flush()
.await
.map_err(PgError::Transport)?;
let mut rows = Vec::new();
loop {
let msg = self
.conn
.codec
.read_message(&mut self.conn.transport)
.await?;
if self.conn.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::RowDescription(body) => {
self.columns = Arc::new(read_row_description(body)?);
}
BackendMessage::DataRow(body) => {
let values = read_data_row(body)?;
rows.push(Row::new(self.columns.clone(), values));
}
BackendMessage::CommandComplete(_body) => {
self.done = true;
}
BackendMessage::PortalSuspended => {
}
BackendMessage::ReadyForQuery(body) => {
self.conn.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.conn.state = ConnectionState::Idle;
break;
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.conn.read_until_ready().await?;
self.conn.state = ConnectionState::Idle;
return Err(PgError::Server(Box::new(server_err)));
}
_ => {}
}
}
Ok(rows)
}
#[must_use = "cursor close errors should be checked"]
pub async fn close(mut self) -> Result<()> {
self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
self.conn
.codec
.encode_and_write(
&mut self.conn.transport,
&FrontendMessage::Close {
variant: b'P',
name: self.portal_name.clone(),
},
)
.await?;
self.conn
.codec
.encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
.await?;
self.conn
.transport
.flush()
.await
.map_err(PgError::Transport)?;
loop {
let msg = self
.conn
.codec
.read_message(&mut self.conn.transport)
.await?;
if self.conn.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::CloseComplete => {}
BackendMessage::ReadyForQuery(body) => {
self.conn.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.conn.state = ConnectionState::Idle;
break;
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.conn.read_until_ready().await?;
self.conn.state = ConnectionState::Idle;
return Err(PgError::Server(Box::new(server_err)));
}
_ => {}
}
}
if self.owns_transaction {
self.conn.execute("COMMIT").await?;
}
self.done = true;
Ok(())
}
pub fn is_done(&self) -> bool {
self.done
}
}
#[derive(Debug)]
enum CursorStreamState {
Active,
Done { command_tag: CommandTag },
Error,
}
#[non_exhaustive]
pub struct CursorStream<'a> {
conn: &'a mut Connection,
portal_name: String,
columns: Arc<Vec<FieldDescription>>,
fetch_size: i32,
state: CursorStreamState,
buffered_rows: Vec<Row>,
owns_transaction: bool,
}
impl<'a> CursorStream<'a> {
pub(crate) fn new(
conn: &'a mut Connection,
portal_name: String,
columns: Arc<Vec<FieldDescription>>,
fetch_size: i32,
owns_transaction: bool,
) -> Self {
CursorStream {
conn,
portal_name,
columns,
fetch_size,
state: CursorStreamState::Active,
buffered_rows: Vec::new(),
owns_transaction,
}
}
#[must_use = "cursor stream errors should be checked"]
pub async fn next(&mut self) -> Result<Option<Row>> {
loop {
match self.state {
CursorStreamState::Done { .. } | CursorStreamState::Error => {
if let Some(row) = self.buffered_rows.pop() {
return Ok(Some(row));
}
return Ok(None);
}
CursorStreamState::Active => {
if let Some(row) = self.buffered_rows.pop() {
return Ok(Some(row));
}
self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
self.conn
.codec
.encode_and_write(
&mut self.conn.transport,
&FrontendMessage::Execute {
portal: self.portal_name.clone(),
max_rows: self.fetch_size,
},
)
.await?;
self.conn
.codec
.encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
.await?;
self.conn
.transport
.flush()
.await
.map_err(PgError::Transport)?;
let mut command_tag: Option<CommandTag> = None;
let mut rows: Vec<Row> = Vec::new();
loop {
let msg = self
.conn
.codec
.read_message(&mut self.conn.transport)
.await?;
if self.conn.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::RowDescription(body) => {
self.columns = Arc::new(read_row_description(body)?);
}
BackendMessage::DataRow(body) => {
let values = read_data_row(body)?;
rows.push(Row::new(self.columns.clone(), values));
}
BackendMessage::CommandComplete(body) => {
command_tag =
Some(CommandTag::new(body.tag().unwrap_or("").into()));
}
BackendMessage::PortalSuspended => {
}
BackendMessage::ReadyForQuery(body) => {
self.conn.transaction_status =
TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.conn.state = ConnectionState::Idle;
break;
}
BackendMessage::ErrorResponse(body) => {
let server_err =
PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.conn.read_until_ready().await?;
self.conn.state = ConnectionState::Idle;
self.state = CursorStreamState::Error;
return Err(PgError::Server(Box::new(server_err)));
}
_ => {}
}
}
if let Some(tag) = command_tag {
self.state = CursorStreamState::Done { command_tag: tag };
}
if rows.is_empty() && !self.is_done() {
continue;
}
rows.reverse();
self.buffered_rows = rows;
if let Some(row) = self.buffered_rows.pop() {
return Ok(Some(row));
}
return Ok(None);
}
}
}
}
pub fn columns(&self) -> &[FieldDescription] {
&self.columns
}
pub fn is_done(&self) -> bool {
matches!(
self.state,
CursorStreamState::Done { .. } | CursorStreamState::Error
)
}
pub fn command_tag(&self) -> Option<&CommandTag> {
match &self.state {
CursorStreamState::Done { command_tag } => Some(command_tag),
_ => None,
}
}
#[must_use = "consume errors should be checked"]
pub async fn consume(mut self) -> Result<CommandTag> {
while self.next().await?.is_some() {}
self.close_portal().await?;
match &self.state {
CursorStreamState::Done { command_tag } => Ok(command_tag.clone()),
_ => Ok(CommandTag::default()),
}
}
async fn close_portal(&mut self) -> Result<()> {
if matches!(self.state, CursorStreamState::Done { .. }) {
if self.owns_transaction {
self.conn.execute("COMMIT").await?;
}
return Ok(());
}
self.conn.transition(ConnectionState::ActiveExtendedQuery)?;
self.conn
.codec
.encode_and_write(
&mut self.conn.transport,
&FrontendMessage::Close {
variant: b'P',
name: self.portal_name.clone(),
},
)
.await?;
self.conn
.codec
.encode_and_write(&mut self.conn.transport, &FrontendMessage::Sync)
.await?;
self.conn
.transport
.flush()
.await
.map_err(PgError::Transport)?;
loop {
let msg = self
.conn
.codec
.read_message(&mut self.conn.transport)
.await?;
if self.conn.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::CloseComplete => {}
BackendMessage::ReadyForQuery(body) => {
self.conn.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.conn.state = ConnectionState::Idle;
break;
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.conn.read_until_ready().await?;
self.conn.state = ConnectionState::Idle;
return Err(PgError::Server(Box::new(server_err)));
}
_ => {}
}
}
if self.owns_transaction {
self.conn.execute("COMMIT").await?;
}
self.state = CursorStreamState::Done {
command_tag: CommandTag::default(),
};
Ok(())
}
}
impl<'a> Drop for CursorStream<'a> {
fn drop(&mut self) {
if !self.is_done() {
self.conn.needs_recovery = true;
}
}
}
impl Connection {
#[must_use = "cursor errors should be checked"]
pub async fn query_cursor(
&mut self,
sql: &str,
params: &[&dyn crate::types::ToSql],
fetch_size: i32,
) -> Result<Cursor<'_>> {
let need_transaction = self.transaction_status == crate::protocol::TransactionStatus::Idle;
if need_transaction {
self.query("BEGIN").await?;
}
self.transition(ConnectionState::ActiveExtendedQuery)?;
let param_values = encode_params_text(params)?;
let portal_name = format!("__pg_portal_{}", self.statement_counter);
self.statement_counter += 1;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Parse {
name: String::new(),
sql: sql.to_string(),
param_types: vec![],
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Bind {
portal: portal_name.clone(),
statement: String::new(),
param_formats: vec![crate::protocol::FormatCode::Text],
params: param_values,
result_formats: vec![crate::protocol::FormatCode::Binary],
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Describe {
variant: b'P',
name: portal_name.clone(),
},
)
.await?;
self.codec
.encode_and_write(&mut self.transport, &FrontendMessage::Sync)
.await?;
self.transport.flush().await.map_err(PgError::Transport)?;
let mut columns: Option<Arc<Vec<FieldDescription>>> = None;
loop {
let msg = self.codec.read_message(&mut self.transport).await?;
if self.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::ParseComplete => {}
BackendMessage::BindComplete => {}
BackendMessage::NoData => {
}
BackendMessage::RowDescription(body) => {
columns = Some(Arc::new(read_row_description(body)?));
}
BackendMessage::ReadyForQuery(body) => {
self.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.state = ConnectionState::Idle;
break;
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.read_until_ready().await?;
self.state = ConnectionState::Idle;
return Err(PgError::Server(Box::new(server_err)));
}
_ => {}
}
}
Ok(Cursor {
conn: self,
portal_name,
columns: columns.unwrap_or_default(),
fetch_size,
done: false,
owns_transaction: need_transaction,
})
}
#[must_use = "cursor stream errors should be checked"]
pub async fn query_cursor_stream(
&mut self,
sql: &str,
params: &[&dyn crate::types::ToSql],
fetch_size: i32,
) -> Result<CursorStream<'_>> {
let need_transaction = self.transaction_status == crate::protocol::TransactionStatus::Idle;
if need_transaction {
self.query("BEGIN").await?;
}
self.transition(ConnectionState::ActiveExtendedQuery)?;
let param_values = encode_params_text(params)?;
let portal_name = format!("__pg_portal_{}", self.statement_counter);
self.statement_counter += 1;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Parse {
name: String::new(),
sql: sql.to_string(),
param_types: vec![],
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Bind {
portal: portal_name.clone(),
statement: String::new(),
param_formats: vec![crate::protocol::FormatCode::Text],
params: param_values,
result_formats: vec![crate::protocol::FormatCode::Binary],
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Describe {
variant: b'P',
name: portal_name.clone(),
},
)
.await?;
self.codec
.encode_and_write(&mut self.transport, &FrontendMessage::Sync)
.await?;
self.transport.flush().await.map_err(PgError::Transport)?;
let mut columns: Option<Arc<Vec<FieldDescription>>> = None;
loop {
let msg = self.codec.read_message(&mut self.transport).await?;
if self.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::ParseComplete => {}
BackendMessage::BindComplete => {}
BackendMessage::NoData => {
}
BackendMessage::RowDescription(body) => {
columns = Some(Arc::new(read_row_description(body)?));
}
BackendMessage::ReadyForQuery(body) => {
self.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.state = ConnectionState::Idle;
break;
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.read_until_ready().await?;
self.state = ConnectionState::Idle;
return Err(PgError::Server(Box::new(server_err)));
}
_ => {}
}
}
Ok(CursorStream::new(
self,
portal_name,
columns.unwrap_or_default(),
fetch_size,
need_transaction,
))
}
}