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, QueryResult};
use crate::query::row::{FieldDescription, Row};
use crate::query::{read_data_row, read_row_description};
use crate::transport::AsyncTransport;
#[derive(Debug)]
pub(crate) enum PipelineOp {
Query {
sql: String,
params: Vec<Option<Vec<u8>>>,
},
Execute {
sql: String,
params: Vec<Option<Vec<u8>>>,
},
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum PipelineResult {
Query(QueryResult),
Execute(CommandTag),
}
#[non_exhaustive]
pub struct Pipeline<'a> {
conn: &'a mut Connection,
ops: Vec<PipelineOp>,
}
impl<'a> Pipeline<'a> {
pub(crate) fn new(conn: &'a mut Connection) -> Self {
Self {
conn,
ops: Vec::new(),
}
}
pub fn query(mut self, sql: &str, params: &[&dyn crate::types::ToSql]) -> Result<Self> {
let values = encode_params_text(params)?;
self.ops.push(PipelineOp::Query {
sql: sql.to_string(),
params: values,
});
Ok(self)
}
pub fn execute(mut self, sql: &str, params: &[&dyn crate::types::ToSql]) -> Result<Self> {
let values = encode_params_text(params)?;
self.ops.push(PipelineOp::Execute {
sql: sql.to_string(),
params: values,
});
Ok(self)
}
#[must_use = "pipeline errors should be checked"]
pub async fn finish(self) -> Result<Vec<PipelineResult>> {
let conn = self.conn;
conn.transition(ConnectionState::ActiveExtendedQuery)?;
for op in &self.ops {
match op {
PipelineOp::Query { sql, params } | PipelineOp::Execute { sql, params } => {
conn.codec
.encode_and_write(
&mut conn.transport,
&FrontendMessage::Parse {
name: String::new(),
sql: sql.clone(),
param_types: vec![],
},
)
.await?;
conn.codec
.encode_and_write(
&mut conn.transport,
&FrontendMessage::Bind {
portal: String::new(),
statement: String::new(),
param_formats: vec![crate::protocol::FormatCode::Text],
params: params.clone(),
result_formats: vec![crate::protocol::FormatCode::Binary],
},
)
.await?;
conn.codec
.encode_and_write(
&mut conn.transport,
&FrontendMessage::Describe {
variant: b'P',
name: String::new(),
},
)
.await?;
conn.codec
.encode_and_write(
&mut conn.transport,
&FrontendMessage::Execute {
portal: String::new(),
max_rows: 0,
},
)
.await?;
}
}
}
conn.codec
.encode_and_write(&mut conn.transport, &FrontendMessage::Sync)
.await?;
conn.transport.flush().await.map_err(PgError::Transport)?;
let mut results = Vec::with_capacity(self.ops.len());
let mut current_op = 0;
let mut current_columns: Option<Arc<Vec<FieldDescription>>> = None;
let mut current_rows: Vec<Row> = Vec::new();
loop {
let msg = conn.codec.read_message(&mut conn.transport).await?;
if conn.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::ParseComplete => {}
BackendMessage::BindComplete => {}
BackendMessage::NoData => {
}
BackendMessage::RowDescription(body) => {
current_columns = Some(Arc::new(read_row_description(body)?));
}
BackendMessage::DataRow(body) => {
let values = read_data_row(body)?;
current_rows.push(Row::new(
current_columns.clone().unwrap_or_default(),
values,
));
}
BackendMessage::CommandComplete(body) => {
let tag = CommandTag::new(body.tag().unwrap_or("").into());
match self.ops.get(current_op) {
Some(PipelineOp::Query { .. }) => {
results.push(PipelineResult::Query(QueryResult::new(
std::mem::take(&mut current_rows),
tag,
current_columns.take().unwrap_or_default(),
)));
}
Some(PipelineOp::Execute { .. }) => {
results.push(PipelineResult::Execute(tag));
}
None => {}
}
current_op += 1;
}
BackendMessage::EmptyQueryResponse => {
match self.ops.get(current_op) {
Some(PipelineOp::Query { .. }) => {
results.push(PipelineResult::Query(QueryResult::new(
Vec::new(),
CommandTag::new("".into()),
Arc::new(Vec::new()),
)));
}
Some(PipelineOp::Execute { .. }) => {
results.push(PipelineResult::Execute(CommandTag::new("".into())));
}
None => {}
}
current_op += 1;
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
conn.read_until_ready().await?;
conn.state = ConnectionState::Idle;
return Err(PgError::Server(Box::new(server_err)));
}
BackendMessage::ReadyForQuery(body) => {
conn.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
conn.state = ConnectionState::Idle;
break;
}
_ => {}
}
}
Ok(results)
}
}
impl Connection {
pub fn pipeline(&mut self) -> Pipeline<'_> {
Pipeline::new(self)
}
}