use std::sync::Arc;
use crate::protocol::{BackendMessage, FrontendMessage, TransactionStatus};
use crate::types::{Format, ToSql};
use crate::connection::{Connection, ConnectionState};
use crate::error::{PgError, PgServerError, Result};
use crate::query::prepared::PreparedStatement;
use crate::query::result::{CommandTag, ExecuteResult, QueryResult};
use crate::query::row::{FieldDescription, Row};
use crate::query::{read_data_row, read_row_description};
use crate::transport::AsyncTransport;
pub(crate) fn encode_params_text(params: &[&dyn ToSql]) -> Result<Vec<Option<Vec<u8>>>> {
let mut values = Vec::with_capacity(params.len());
for p in params {
let mut buf = Vec::new();
let is_null = p.to_sql(&crate::types::Type::UNKNOWN, &mut buf, Format::Text)?;
match is_null {
crate::types::IsNull::Yes => values.push(None),
crate::types::IsNull::No => values.push(Some(buf)),
}
}
Ok(values)
}
pub(crate) fn encode_params_binary(
params: &[&dyn ToSql],
param_types: &[crate::types::Type],
) -> Result<Vec<Option<Vec<u8>>>> {
if params.len() != param_types.len() {
return Err(PgError::Config(format!(
"parameter count mismatch: expected {}, got {}",
param_types.len(),
params.len()
)));
}
let mut values = Vec::with_capacity(params.len());
for (p, ty) in params.iter().zip(param_types.iter()) {
let mut buf = Vec::new();
let is_null = p.to_sql(ty, &mut buf, Format::Binary)?;
match is_null {
crate::types::IsNull::Yes => values.push(None),
crate::types::IsNull::No => values.push(Some(buf)),
}
}
Ok(values)
}
impl Connection {
#[must_use = "query results should be checked for errors"]
pub async fn query_params(&mut self, sql: &str, params: &[&dyn ToSql]) -> Result<QueryResult> {
self.transition(ConnectionState::ActiveExtendedQuery)?;
let param_values = encode_params_text(params)?;
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: String::new(),
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: String::new(),
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Execute {
portal: String::new(),
max_rows: 0,
},
)
.await?;
self.codec
.encode_and_write(&mut self.transport, &FrontendMessage::Sync)
.await?;
self.transport.flush().await.map_err(PgError::Transport)?;
let result = self.read_extended_query_result().await;
self.state = ConnectionState::Idle;
result
}
#[must_use = "execute results should be checked for errors"]
pub async fn execute_params(
&mut self,
sql: &str,
params: &[&dyn ToSql],
) -> Result<ExecuteResult> {
let result = self.query_params(sql, params).await?;
Ok(ExecuteResult::new(result.command_tag().clone()))
}
#[must_use = "stream errors should be checked"]
pub async fn query_params_stream(
&mut self,
sql: &str,
params: &[&dyn ToSql],
) -> Result<crate::query::stream::RowStream<'_>> {
use crate::query::stream::RowStream;
#[cfg(feature = "tracing")]
tracing::debug!(
target: crate::tracing_ext::TARGET_QUERY,
sql_len = sql.len(),
sql_truncated = %crate::tracing_ext::truncate_str(sql, 200),
param_count = params.len(),
protocol = "extended",
"Executing parameterized query"
);
#[cfg(feature = "tracing")]
tracing::trace!(target: crate::tracing_ext::TARGET_QUERY, sql = %sql, "Full SQL text");
if self.needs_recovery {
self.recover().await?;
}
self.transition(ConnectionState::Streaming)?;
let param_values = encode_params_text(params)?;
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: String::new(),
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: String::new(),
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Execute {
portal: String::new(),
max_rows: 0,
},
)
.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::RowDescription(body) => {
columns = Some(Arc::new(read_row_description(body)?));
break;
}
BackendMessage::NoData => {
break;
}
BackendMessage::CommandComplete(_body) => {
break;
}
BackendMessage::ReadyForQuery(body) => {
self.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::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)));
}
BackendMessage::EmptyQueryResponse => {
break;
}
_ => {}
}
}
Ok(RowStream::new_extended_receiving(self, columns))
}
#[must_use = "query results should be checked for errors"]
pub async fn query_prepared(
&mut self,
stmt: &PreparedStatement,
params: &[&dyn ToSql],
) -> Result<QueryResult> {
self.transition(ConnectionState::ActiveExtendedQuery)?;
let param_values = encode_params_binary(params, &stmt.param_types)?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Bind {
portal: String::new(),
statement: stmt.name.clone(),
param_formats: vec![crate::protocol::FormatCode::Binary],
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: String::new(),
},
)
.await?;
self.codec
.encode_and_write(
&mut self.transport,
&FrontendMessage::Execute {
portal: String::new(),
max_rows: 0,
},
)
.await?;
self.codec
.encode_and_write(&mut self.transport, &FrontendMessage::Sync)
.await?;
self.transport.flush().await.map_err(PgError::Transport)?;
let result = self.read_extended_query_result().await;
self.state = ConnectionState::Idle;
result
}
#[must_use = "execute results should be checked for errors"]
pub async fn execute_prepared(
&mut self,
stmt: &PreparedStatement,
params: &[&dyn ToSql],
) -> Result<ExecuteResult> {
let result = self.query_prepared(stmt, params).await?;
Ok(ExecuteResult::new(result.command_tag().clone()))
}
}
impl Connection {
async fn read_extended_query_result(&mut self) -> Result<QueryResult> {
let mut columns: Option<Arc<Vec<FieldDescription>>> = None;
let mut rows: Vec<Row> = Vec::new();
let mut tag = 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::DataRow(body) => {
let values = read_data_row(body)?;
if columns.is_none() {
let synthetic: Vec<FieldDescription> = (0..values.len())
.map(|i| {
FieldDescription::new(
format!("col{}", i),
0,
0,
0, -1,
-1,
1, )
})
.collect();
columns = Some(Arc::new(synthetic));
}
rows.push(Row::new(columns.clone().unwrap(), values));
}
BackendMessage::CommandComplete(body) => {
tag = Some(CommandTag::new(body.tag().unwrap_or("").into()));
}
BackendMessage::EmptyQueryResponse => {
tag = Some(CommandTag::new("".into()));
}
BackendMessage::ErrorResponse(body) => {
let server_err = PgServerError::from_error_body(&body).map_err(PgError::Io)?;
self.read_until_ready().await?;
return Err(PgError::Server(Box::new(server_err)));
}
BackendMessage::ReadyForQuery(body) => {
self.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
break;
}
_ => {}
}
}
Ok(QueryResult::new(
rows,
tag.unwrap_or_default(),
columns.unwrap_or_default(),
))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::auth::{Codec, ServerParams};
use crate::config::Config;
use crate::connection::ConnectionState;
use crate::transport::{BufferedTransport, ClientTransport, MockTransport, PgTransport};
use std::collections::VecDeque;
fn make_connection(read_data: Vec<u8>) -> Connection {
let transport = PgTransport::Plain(BufferedTransport::new(ClientTransport::Mock(
MockTransport::new(read_data),
)));
Connection {
transport,
codec: Codec::new(),
server_params: ServerParams::default(),
state: ConnectionState::Idle,
config: Config::new(),
transaction_status: TransactionStatus::Idle,
notification_queue: VecDeque::new(),
notice_handler: None,
statement_counter: 0,
needs_recovery: false,
health: crate::reconnect::session::ConnectionHealth::new(),
session_state: crate::reconnect::session::SessionState::new(),
}
}
fn build_parse_complete() -> Vec<u8> {
vec![b'1', 0, 0, 0, 4]
}
fn build_bind_complete() -> Vec<u8> {
vec![b'2', 0, 0, 0, 4]
}
fn build_row_description_msg(fields: &[(&str, u32)]) -> Vec<u8> {
let mut buf = vec![b'T'];
let mut body = Vec::new();
body.extend_from_slice(&(fields.len() as i16).to_be_bytes());
for (name, type_oid) in fields {
body.extend_from_slice(name.as_bytes());
body.push(0);
body.extend_from_slice(&0u32.to_be_bytes()); body.extend_from_slice(&0i16.to_be_bytes()); body.extend_from_slice(&type_oid.to_be_bytes()); body.extend_from_slice(&(-1i16).to_be_bytes()); body.extend_from_slice(&(-1i32).to_be_bytes()); body.extend_from_slice(&0i16.to_be_bytes()); }
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
fn build_data_row_msg(values: &[Option<&str>]) -> Vec<u8> {
let mut buf = vec![b'D'];
let mut body = Vec::new();
body.extend_from_slice(&(values.len() as i16).to_be_bytes());
for val in values {
match val {
Some(v) => {
let bytes = v.as_bytes();
body.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
body.extend_from_slice(bytes);
}
None => {
body.extend_from_slice(&(-1i32).to_be_bytes());
}
}
}
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
fn build_command_complete_msg(tag: &str) -> Vec<u8> {
let mut buf = vec![b'C'];
let mut body = Vec::new();
body.extend_from_slice(tag.as_bytes());
body.push(0);
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
fn build_ready_for_query(status: u8) -> Vec<u8> {
vec![b'Z', 0, 0, 0, 5, status]
}
#[tokio::test]
async fn test_query_params_select() {
let mut data = Vec::new();
data.extend_from_slice(&build_parse_complete());
data.extend_from_slice(&build_bind_complete());
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("42")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn.query_params("SELECT $1", &[&42i32]).await.unwrap();
assert_eq!(result.len(), 1);
let v: i32 = result.rows()[0].get(0).unwrap();
assert_eq!(v, 42);
}
#[tokio::test]
async fn test_query_params_insert() {
let mut data = Vec::new();
data.extend_from_slice(&build_parse_complete());
data.extend_from_slice(&build_bind_complete());
data.extend_from_slice(&build_command_complete_msg("INSERT 0 1"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.execute_params("INSERT INTO t (id) VALUES ($1)", &[&1i32])
.await
.unwrap();
assert_eq!(result.rows_affected(), Some(1));
}
#[tokio::test]
async fn test_query_params_insert_two_params() {
let mut data = Vec::new();
data.extend_from_slice(&build_parse_complete());
data.extend_from_slice(&build_bind_complete());
data.extend_from_slice(&build_command_complete_msg("INSERT 0 1"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.execute_params(
"INSERT INTO t (id, name) VALUES ($1, $2)",
&[&200i32, &"param_insert"],
)
.await
.unwrap();
assert_eq!(result.rows_affected(), Some(1));
}
#[tokio::test]
async fn test_query_prepared_select() {
let mut data = Vec::new();
data.extend_from_slice(&build_bind_complete());
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("99")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let stmt = PreparedStatement {
name: "__pg_stmt_1".into(),
sql: "SELECT $1".into(),
param_types: vec![crate::types::Type::INT4],
columns: Arc::new(vec![]),
};
let result = conn.query_prepared(&stmt, &[&99i32]).await.unwrap();
assert_eq!(result.len(), 1);
let v: i32 = result.rows()[0].get(0).unwrap();
assert_eq!(v, 99);
}
#[tokio::test]
async fn test_query_params_error() {
let mut data = Vec::new();
let mut err = vec![b'E', 0, 0, 0, 22];
err.extend_from_slice(b"S");
err.extend_from_slice(b"ERROR\0");
err.extend_from_slice(b"M");
err.extend_from_slice(b"syntax error\0");
err.push(0);
data.extend_from_slice(&err);
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn.query_params("BAD $1", &[&1i32]).await;
assert!(result.is_err());
assert!(conn.is_idle());
}
fn build_error_response_msg(code: &str, message: &str) -> Vec<u8> {
let mut buf = vec![b'E'];
let mut body = Vec::new();
body.push(b'S');
body.extend_from_slice(b"ERROR");
body.push(0);
body.push(b'C');
body.extend_from_slice(code.as_bytes());
body.push(0);
body.push(b'M');
body.extend_from_slice(message.as_bytes());
body.push(0);
body.push(0);
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
#[tokio::test]
async fn test_query_params_stream() {
let mut data = Vec::new();
data.extend_from_slice(&build_parse_complete());
data.extend_from_slice(&build_bind_complete());
data.extend_from_slice(&build_row_description_msg(&[
("id", crate::types::INT4_OID),
("name", crate::types::TEXT_OID),
]));
data.extend_from_slice(&build_data_row_msg(&[Some("1"), Some("Alice")]));
data.extend_from_slice(&build_data_row_msg(&[Some("2"), Some("Bob")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 2"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let mut stream = conn
.query_params_stream("SELECT * FROM users WHERE id > $1", &[&0i32])
.await
.unwrap();
let row1 = stream.next().await.unwrap().unwrap();
assert_eq!(row1.get::<i32>(0).unwrap(), 1);
assert_eq!(row1.get::<String>(1).unwrap(), "Alice");
let row2 = stream.next().await.unwrap().unwrap();
assert_eq!(row2.get::<i32>(0).unwrap(), 2);
assert_eq!(row2.get::<String>(1).unwrap(), "Bob");
assert!(stream.next().await.unwrap().is_none());
assert!(stream.is_done());
}
#[tokio::test]
async fn test_query_params_stream_error() {
let mut data = Vec::new();
data.extend_from_slice(&build_parse_complete());
data.extend_from_slice(&build_bind_complete());
data.extend_from_slice(&build_error_response_msg("23505", "duplicate key value"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.query_params_stream("INSERT INTO t VALUES ($1)", &[&1i32])
.await;
assert!(result.is_err());
match result {
Err(PgError::Server(e)) => {
assert_eq!(e.code(), "23505");
}
_ => panic!("expected server error"),
}
}
#[tokio::test]
async fn test_query_params_stream_no_data() {
let mut data = Vec::new();
data.extend_from_slice(&build_parse_complete());
data.extend_from_slice(&build_bind_complete());
let nodata = vec![b'n', 0, 0, 0, 4];
data.extend_from_slice(&nodata);
data.extend_from_slice(&build_command_complete_msg("INSERT 0 1"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let mut stream = conn
.query_params_stream("INSERT INTO t VALUES ($1)", &[&1i32])
.await
.unwrap();
assert!(stream.next().await.unwrap().is_none());
assert!(stream.is_done());
}
}