use std::sync::Arc;
use crate::protocol::{BackendMessage, FrontendMessage, TransactionStatus};
use fallible_iterator::FallibleIterator;
use crate::connection::Connection;
use crate::error::{PgError, PgServerError, Result};
use crate::query::result::{CommandTag, ExecuteResult, QueryResult};
use crate::query::row::{FieldDescription, Row};
use crate::transport::AsyncTransport;
#[cfg(feature = "tracing")]
use crate::tracing_ext::{truncate_str, TARGET_QUERY};
pub mod cache;
pub mod cursor;
pub mod params;
pub mod pipeline;
pub mod prepared;
pub mod result;
pub mod row;
pub mod stream;
pub use cache::StatementCache;
pub use cursor::Cursor;
pub use cursor::CursorStream;
pub use pipeline::{Pipeline, PipelineResult};
pub use prepared::PreparedStatement;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Notice {
inner: PgServerError,
}
pub type NoticeHandler = Box<dyn Fn(&Notice) + Send + Sync>;
impl Notice {
pub fn from_fields(fields: &crate::protocol::backend::NoticeResponseBody) -> Result<Self> {
let inner = PgServerError::from_notice_body(fields).map_err(PgError::Io)?;
Ok(Self { inner })
}
pub fn severity(&self) -> &str {
&self.inner.severity
}
pub fn code(&self) -> &str {
&self.inner.code
}
pub fn message(&self) -> &str {
&self.inner.message
}
pub fn detail(&self) -> Option<&str> {
self.inner.detail.as_deref()
}
pub fn hint(&self) -> Option<&str> {
self.inner.hint.as_deref()
}
pub fn as_server_error(&self) -> &PgServerError {
&self.inner
}
}
impl std::fmt::Display for Notice {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"{}: {} (SQLSTATE {})",
self.inner.severity, self.inner.message, self.inner.code
)?;
if let Some(detail) = &self.inner.detail {
write!(f, "\nDETAIL: {}", detail)?;
}
if let Some(hint) = &self.inner.hint {
write!(f, "\nHINT: {}", hint)?;
}
Ok(())
}
}
impl Connection {
#[must_use = "query results should be checked for errors"]
pub async fn query(&mut self, sql: &str) -> Result<QueryResult> {
let mut stream = self.query_stream(sql).await?;
let mut rows = Vec::new();
while let Some(row) = stream.next().await? {
rows.push(row);
}
let columns = stream.columns().map(|c| c.to_vec()).unwrap_or_default();
let command_tag = stream.command_tag().cloned().unwrap_or_default();
Ok(QueryResult::new(rows, command_tag, Arc::new(columns)))
}
#[must_use = "execute results should be checked for errors"]
pub async fn execute(&mut self, sql: &str) -> Result<ExecuteResult> {
let result = self.query(sql).await?;
Ok(ExecuteResult::new(result.command_tag().clone()))
}
#[must_use = "query results should be checked for errors"]
pub async fn query_one(&mut self, sql: &str) -> Result<Option<Row>> {
let result = self.query(sql).await?;
Ok(result.into_rows().into_iter().next())
}
#[must_use = "query results should be checked for errors"]
pub async fn query_each<F>(&mut self, sql: &str, mut f: F) -> Result<CommandTag>
where
F: FnMut(Row) -> Result<()>,
{
let mut stream = self.query_stream(sql).await?;
while let Some(row) = stream.next().await? {
f(row)?;
}
stream
.command_tag()
.cloned()
.ok_or_else(|| PgError::InvalidState("stream ended without command tag".into()))
}
#[must_use = "batch results should be checked for errors"]
pub async fn batch_execute(&mut self, sql: &str) -> Result<Vec<QueryResult>> {
self.transition(ConnectionState::ActiveSimpleQuery)?;
self.codec
.send(
&mut self.transport,
&FrontendMessage::Query { sql: sql.into() },
)
.await?;
let mut results = Vec::new();
let mut current_columns: Option<Arc<Vec<FieldDescription>>> = None;
let mut current_rows: Vec<Row> = Vec::new();
loop {
let msg = self.codec.read_message(&mut self.transport).await?;
if self.handle_async_message(&msg) {
continue;
}
match msg {
BackendMessage::RowDescription(body) => {
current_columns = Some(Arc::new(read_row_description(body)?));
current_rows.clear();
}
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());
results.push(QueryResult::new(
std::mem::take(&mut current_rows),
tag,
current_columns.take().unwrap_or_default(),
));
}
BackendMessage::EmptyQueryResponse => {
results.push(QueryResult::new(
Vec::new(),
CommandTag::new("".into()),
Arc::new(Vec::new()),
));
}
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::ReadyForQuery(body) => {
self.transaction_status = TransactionStatus::from_u8(body.status())
.unwrap_or(TransactionStatus::Idle);
self.state = ConnectionState::Idle;
break;
}
_ => {}
}
}
Ok(results)
}
}
impl Connection {
#[must_use = "stream results should be checked for errors"]
pub async fn query_stream(&mut self, sql: &str) -> Result<stream::RowStream<'_>> {
#[cfg(feature = "tracing")]
tracing::debug!(target: TARGET_QUERY, sql_len = sql.len(), sql_truncated = %truncate_str(sql, 200), protocol = "simple", "Executing simple query");
#[cfg(feature = "tracing")]
tracing::trace!(target: TARGET_QUERY, sql = %sql, "Full SQL text");
if self.needs_recovery {
self.recover().await?;
}
self.transition(ConnectionState::Streaming)?;
self.codec
.send(
&mut self.transport,
&FrontendMessage::Query { sql: sql.into() },
)
.await?;
Ok(stream::RowStream::new_simple(self))
}
#[must_use = "stream results should be checked for errors"]
pub async fn query_prepared_stream(
&mut self,
stmt: &PreparedStatement,
params: &[&dyn crate::types::ToSql],
) -> Result<stream::RowStream<'_>> {
#[cfg(feature = "tracing")]
tracing::debug!(target: TARGET_QUERY, sql_len = stmt.sql.len(), sql_truncated = %truncate_str(&stmt.sql, 200), statement = %stmt.name, "Executing prepared statement");
if self.needs_recovery {
self.recover().await?;
}
self.transition(ConnectionState::Streaming)?;
let param_values = params::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)?;
Ok(stream::RowStream::new_extended_with_columns(
self,
stmt.columns.clone(),
))
}
#[must_use = "query results should be checked for errors"]
pub async fn query_each_async<F, Fut>(&mut self, sql: &str, mut f: F) -> Result<CommandTag>
where
F: FnMut(Row) -> Fut,
Fut: std::future::Future<Output = Result<()>>,
{
let mut stream = self.query_stream(sql).await?;
while let Some(row) = stream.next().await? {
f(row).await?;
}
stream
.command_tag()
.cloned()
.ok_or_else(|| PgError::InvalidState("stream ended without command tag".into()))
}
}
use crate::connection::ConnectionState;
pub(crate) fn read_row_description(
body: crate::protocol::backend::RowDescriptionBody,
) -> Result<Vec<FieldDescription>> {
let mut fields = Vec::new();
let mut iter = body.fields();
while let Some(field) = iter.next()? {
fields.push(FieldDescription::new(
field.name().into(),
field.table_oid(),
field.column_id(),
field.type_oid(),
field.type_size(),
field.type_modifier(),
field.format(),
));
}
Ok(fields)
}
pub(crate) fn read_data_row(
body: crate::protocol::backend::DataRowBody,
) -> Result<Vec<Option<Vec<u8>>>> {
let buf = body.buffer();
let mut values = Vec::new();
let mut iter = body.ranges();
while let Some(range) = iter.next()? {
values.push(range.map(|r| buf[r].to_vec()));
}
Ok(values)
}
#[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(),
}
}
pub(crate) 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_basic() {
let mut data = Vec::new();
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 result = conn.query("SELECT id, name FROM users").await.unwrap();
assert_eq!(result.len(), 2);
let id: i32 = result.rows()[0].get(0).unwrap();
assert_eq!(id, 1);
let name: String = result.rows()[0].get(1).unwrap();
assert_eq!(name, "alice");
}
#[tokio::test]
async fn test_execute_no_rows() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("INSERT 0 3"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.execute("INSERT INTO users (name) VALUES ('alice')")
.await
.unwrap();
assert_eq!(result.rows_affected(), Some(3));
}
#[tokio::test]
async fn test_query_one() {
let mut data = Vec::new();
data.extend_from_slice(&build_row_description_msg(&[(
"id",
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 row = conn.query_one("SELECT 42").await.unwrap();
assert!(row.is_some());
let id: i32 = row.unwrap().get(0).unwrap();
assert_eq!(id, 42);
}
#[tokio::test]
async fn test_query_empty() {
let mut data = Vec::new();
data.extend_from_slice(&build_row_description_msg(&[(
"id",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_command_complete_msg("SELECT 0"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.query("SELECT id FROM users WHERE false")
.await
.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn test_query_error() {
let mut data = Vec::new();
let mut err = vec![b'E', 0, 0, 0, 26];
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("BAD SQL").await;
assert!(result.is_err());
assert!(conn.is_idle());
}
#[tokio::test]
async fn test_query_each() {
let mut data = Vec::new();
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("10")]));
data.extend_from_slice(&build_data_row_msg(&[Some("20")]));
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 sum = 0i32;
let tag = conn
.query_each("SELECT val FROM nums", |row| {
let v: i32 = row.get(0)?;
sum += v;
Ok(())
})
.await
.unwrap();
assert_eq!(sum, 30);
assert_eq!(tag.as_str(), "SELECT 2");
}
#[tokio::test]
async fn test_batch_execute() {
let mut data = Vec::new();
data.extend_from_slice(&build_row_description_msg(&[(
"id",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("1")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
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 results = conn
.batch_execute("SELECT 1; INSERT INTO t VALUES (1)")
.await
.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].len(), 1);
assert_eq!(results[1].len(), 0);
assert_eq!(results[1].rows_affected(), Some(1));
}
#[tokio::test]
async fn test_null_handling() {
let mut data = Vec::new();
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[None]));
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("SELECT NULL").await.unwrap();
assert_eq!(result.len(), 1);
assert!(result.rows()[0].is_null(0));
}
}