use crate::{EmbeddedDatabase, Error, Result, Schema, Tuple, Value};
#[inline]
pub(super) fn starts_with_icase(s: &str, prefix: &str) -> bool {
#[allow(clippy::indexing_slicing)]
{
s.len() >= prefix.len() && s.as_bytes()[..prefix.len()].eq_ignore_ascii_case(prefix.as_bytes())
}
}
#[inline]
pub(super) fn starts_with_cte(s: &str) -> bool {
let t = s.trim_start();
starts_with_icase(t, "WITH")
&& t.as_bytes()
.get(4)
.map_or(true, |c| !c.is_ascii_alphanumeric() && *c != b'_')
}
use super::auth::{AuthManager, AuthMethod, ScramAuthState};
use super::catalog::PgCatalog;
use super::messages::{AuthenticationMessage, BackendMessage, FieldDescription, FrontendMessage, TransactionStatus};
use super::prepared::PreparedStatementManager;
use super::ssl::SecureConnection;
use bytes::{BufMut, BytesMut};
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt, BufWriter};
use tokio::net::TcpStream;
#[cfg(unix)]
use tokio::net::UnixStream;
pub struct PgConnectionHandler<S = BufWriter<TcpStream>> {
stream: S,
pub(super) database: Arc<EmbeddedDatabase>,
auth_manager: Arc<AuthManager>,
pub(super) catalog: PgCatalog,
pub(super) prepared_statements: PreparedStatementManager,
authenticated: bool,
transaction_status: TransactionStatus,
buffer: BytesMut,
username: Option<String>,
scram_state: Option<ScramAuthState>,
write_buf: BytesMut,
pub(super) suppress_ready_for_query: bool,
awaiting_sync_after_error: bool,
pub(super) session_id: crate::session::SessionId,
}
impl<S> Drop for PgConnectionHandler<S> {
fn drop(&mut self) {
let _ = self.database.destroy_session(self.session_id);
}
}
impl PgConnectionHandler<BufWriter<TcpStream>> {
pub fn new(
stream: TcpStream,
database: Arc<EmbeddedDatabase>,
auth_manager: Arc<AuthManager>,
initial_data: Option<&[u8]>,
) -> Self {
let mut buffer = BytesMut::with_capacity(8192);
if let Some(data) = initial_data {
buffer.extend_from_slice(data);
}
let session_id = database
.create_wire_session("pg_wire")
.expect("wire session creation is infallible");
Self {
stream: BufWriter::new(stream),
database: database.clone(),
auth_manager,
catalog: PgCatalog::with_database(database),
prepared_statements: PreparedStatementManager::new(),
authenticated: false,
transaction_status: TransactionStatus::Idle,
buffer,
username: None,
scram_state: None,
write_buf: BytesMut::with_capacity(4096),
suppress_ready_for_query: false,
awaiting_sync_after_error: false,
session_id,
}
}
}
#[cfg(unix)]
impl PgConnectionHandler<BufWriter<UnixStream>> {
pub fn new_unix(stream: UnixStream, database: Arc<EmbeddedDatabase>, auth_manager: Arc<AuthManager>) -> Self {
let session_id = database
.create_wire_session("pg_wire")
.expect("wire session creation is infallible");
Self {
stream: BufWriter::new(stream),
database: database.clone(),
auth_manager,
catalog: PgCatalog::with_database(database),
prepared_statements: PreparedStatementManager::new(),
authenticated: false,
transaction_status: TransactionStatus::Idle,
buffer: BytesMut::with_capacity(8192),
username: None,
scram_state: None,
write_buf: BytesMut::with_capacity(4096),
suppress_ready_for_query: false,
awaiting_sync_after_error: false,
session_id,
}
}
}
#[cfg(unix)]
pub async fn handle_connection_unix(
database: Arc<EmbeddedDatabase>,
stream: UnixStream,
_connection_id: u32,
) -> Result<()> {
let auth_manager = Arc::new(AuthManager::new(AuthMethod::Trust));
let mut handler = PgConnectionHandler::new_unix(stream, database, auth_manager);
handler.handle().await
}
impl PgConnectionHandler<BufWriter<SecureConnection<TcpStream>>> {
pub fn new_with_stream(
stream: SecureConnection<TcpStream>,
database: Arc<EmbeddedDatabase>,
auth_manager: Arc<AuthManager>,
initial_data: Option<&[u8]>,
) -> Self {
let mut buffer = BytesMut::with_capacity(8192);
if let Some(data) = initial_data {
buffer.extend_from_slice(data);
}
let session_id = database
.create_wire_session("pg_wire")
.expect("wire session creation is infallible");
Self {
stream: BufWriter::new(stream),
database: database.clone(),
auth_manager,
catalog: PgCatalog::with_database(database),
prepared_statements: PreparedStatementManager::new(),
authenticated: false,
transaction_status: TransactionStatus::Idle,
buffer,
username: None,
scram_state: None,
write_buf: BytesMut::with_capacity(4096),
suppress_ready_for_query: false,
awaiting_sync_after_error: false,
session_id,
}
}
}
impl<S> PgConnectionHandler<S>
where
S: AsyncReadExt + AsyncWriteExt + Unpin,
{
#[cfg(test)]
pub(super) fn new_for_tests(database: Arc<EmbeddedDatabase>, stream: S) -> Self {
Self {
stream,
session_id: database
.create_wire_session("pg_wire_test")
.expect("wire session creation is infallible"),
database: Arc::clone(&database),
auth_manager: Arc::new(AuthManager::new(AuthMethod::Trust)),
catalog: PgCatalog::with_database(database),
prepared_statements: PreparedStatementManager::new(),
authenticated: true,
transaction_status: TransactionStatus::Idle,
buffer: BytesMut::with_capacity(8192),
username: None,
scram_state: None,
write_buf: BytesMut::with_capacity(4096),
suppress_ready_for_query: false,
awaiting_sync_after_error: false,
}
}
pub async fn handle(&mut self) -> Result<()> {
tracing::info!("New PostgreSQL connection");
if let Err(e) = self.handle_startup().await {
tracing::error!("Startup failed: {}", e);
let _ = self.send_error("FATAL", "08P01", &e.to_string(), None, None).await;
return Err(e);
}
tracing::debug!("Entering main message loop");
loop {
tracing::trace!("Waiting for next message from client");
match self.read_message().await {
Ok(Some(msg)) => {
tracing::debug!("Received message: {:?}", msg);
if self.awaiting_sync_after_error {
self.handle_message_while_awaiting_sync(msg).await?;
continue;
}
let wait_for_sync = Self::message_requires_sync_after_error(&msg);
if let Err(e) = self.handle_message(msg).await {
tracing::error!("Error handling message: {}", e);
self.send_error_for_query(&e, wait_for_sync).await?;
if wait_for_sync {
self.awaiting_sync_after_error = true;
}
}
}
Ok(None) => {
tracing::info!("Client disconnected");
break;
}
Err(e) => {
tracing::error!("Error reading message: {}", e);
break;
}
}
}
Ok(())
}
async fn handle_message_while_awaiting_sync(&mut self, msg: FrontendMessage) -> Result<()> {
match msg {
FrontendMessage::Sync => {
self.awaiting_sync_after_error = false;
self.send_ready_for_query().await
}
FrontendMessage::Terminate => Ok(()),
_ => {
tracing::debug!("Discarding frontend message until Sync after extended-query error");
Ok(())
}
}
}
fn message_requires_sync_after_error(msg: &FrontendMessage) -> bool {
matches!(
msg,
FrontendMessage::Parse { .. }
| FrontendMessage::Bind { .. }
| FrontendMessage::Execute { .. }
| FrontendMessage::Describe { .. }
| FrontendMessage::Close { .. }
)
}
#[allow(clippy::indexing_slicing)]
async fn handle_startup(&mut self) -> Result<()> {
let len_buf: [u8; 4];
if self.buffer.len() >= 4 {
len_buf = [self.buffer[0], self.buffer[1], self.buffer[2], self.buffer[3]];
} else {
let mut buf = [0u8; 4];
self.stream
.read_exact(&mut buf)
.await
.map_err(|e| Error::network(format!("Failed to read startup length: {}", e)))?;
len_buf = buf;
self.buffer.extend_from_slice(&len_buf);
}
let len = i32::from_be_bytes(len_buf) as usize;
let bytes_in_buffer = self.buffer.len();
let bytes_needed = len.saturating_sub(bytes_in_buffer);
if bytes_needed > 0 {
let mut remaining_buf = vec![0u8; bytes_needed];
self.stream
.read_exact(&mut remaining_buf)
.await
.map_err(|e| Error::network(format!("Failed to read startup message: {}", e)))?;
self.buffer.extend_from_slice(&remaining_buf);
}
let msg = FrontendMessage::parse_startup(&mut self.buffer)?
.ok_or_else(|| Error::protocol("Invalid startup message"))?;
if let FrontendMessage::Startup {
protocol_version,
params,
} = msg
{
tracing::info!("Protocol version: {}, params: {:?}", protocol_version, params);
self.username = params.get("user").cloned();
if let Some(requested) = params.get("database").cloned().or_else(|| params.get("user").cloned()) {
if !self.database.database_name_is_valid(&requested) {
return Err(Error::authentication(format!(
"database \"{requested}\" does not exist"
)));
}
}
match self.auth_manager.method() {
AuthMethod::Trust => {
self.authenticated = true;
self.send_auth_ok().await?;
}
AuthMethod::CleartextPassword => {
self.send_message(BackendMessage::Authentication(AuthenticationMessage::CleartextPassword))
.await?;
self.flush().await?;
if let Some(FrontendMessage::PasswordMessage { password }) = self.read_message().await? {
let username = self
.username
.as_ref()
.ok_or_else(|| Error::authentication("No username provided"))?;
if self.auth_manager.verify_cleartext(username, &password)? {
self.authenticated = true;
self.send_auth_ok().await?;
} else {
return Err(Error::authentication("Invalid password"));
}
} else {
return Err(Error::protocol("Expected password message"));
}
}
AuthMethod::ScramSha256 => {
self.handle_scram_authentication().await?;
}
_ => {
self.authenticated = true;
self.send_auth_ok().await?;
}
}
self.send_parameter_status(
"server_version",
&format!("16.0 (HeliosDB Nano {})", env!("CARGO_PKG_VERSION")),
)
.await?;
self.send_parameter_status("server_encoding", "UTF8").await?;
self.send_parameter_status("client_encoding", "UTF8").await?;
self.send_parameter_status("DateStyle", "ISO, MDY").await?;
self.send_parameter_status("TimeZone", "UTC").await?;
self.send_parameter_status("integer_datetimes", "on").await?;
self.send_parameter_status("helios.copy", "on").await?;
self.send_parameter_status("helios.pipeline", "on").await?;
self.send_parameter_status("helios.plan_cache", "on").await?;
self.send_parameter_status("helios.binary_results", "on").await?;
self.send_parameter_status("helios.reset_session", "on").await?;
self.send_parameter_status("helios.fast_autocommit", "off").await?;
self.send_message(BackendMessage::BackendKeyData {
process_id: std::process::id() as i32,
secret_key: rand::random(),
})
.await?;
self.send_ready_for_query().await?;
Ok(())
} else {
Err(Error::protocol("Expected startup message"))
}
}
#[allow(clippy::indexing_slicing)]
async fn read_message(&mut self) -> Result<Option<FrontendMessage>> {
tracing::trace!("read_message: Checking buffer, len={}", self.buffer.len());
if let Some(msg) = FrontendMessage::parse(&mut self.buffer)? {
tracing::trace!("read_message: Parsed message from existing buffer");
return Ok(Some(msg));
}
let mut temp_buf = vec![0u8; 4096];
loop {
tracing::trace!("read_message: Attempting to read from stream");
match self.stream.read(&mut temp_buf).await {
Ok(0) => {
tracing::debug!("read_message: EOF received (0 bytes)");
return Ok(None); }
Ok(n) => {
tracing::trace!("read_message: Read {} bytes", n);
self.buffer.extend_from_slice(&temp_buf[..n]);
if let Some(msg) = FrontendMessage::parse(&mut self.buffer)? {
tracing::trace!("read_message: Successfully parsed message after read");
return Ok(Some(msg));
}
tracing::trace!("read_message: Insufficient data for complete message, continuing");
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
tracing::trace!("read_message: WouldBlock, sleeping 10ms");
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
}
Err(e) => {
tracing::error!("read_message: Read error: {}", e);
return Err(Error::network(format!("Read error: {}", e)));
}
}
}
}
pub(super) async fn handle_message(&mut self, msg: FrontendMessage) -> Result<()> {
if !self.authenticated && !matches!(msg, FrontendMessage::PasswordMessage { .. }) {
return Err(Error::authentication("Not authenticated"));
}
match msg {
FrontendMessage::Query { query } => {
self.handle_query(&query).await?;
}
FrontendMessage::Parse {
statement_name,
query,
param_types,
} => {
self.handle_parse_extended(statement_name, query, param_types).await?;
}
FrontendMessage::Bind {
portal_name,
statement_name,
param_formats,
params,
result_formats,
} => {
self.handle_bind_extended(portal_name, statement_name, param_formats, params, result_formats)
.await?;
}
FrontendMessage::Execute { portal_name, max_rows } => {
self.handle_execute_extended(portal_name, max_rows).await?;
}
FrontendMessage::Describe { target, name } => {
self.handle_describe_extended(target, name).await?;
}
FrontendMessage::Close { target, name } => {
self.handle_close(target, name).await?;
}
FrontendMessage::Sync => {
self.send_ready_for_query().await?;
}
FrontendMessage::Flush => {
self.flush().await?;
}
FrontendMessage::Terminate => {
return Ok(());
}
_ => {
tracing::warn!("Unhandled message type: {:?}", msg);
}
}
Ok(())
}
async fn handle_query(&mut self, query: &str) -> Result<()> {
let statements = pg_split_sql_respecting_quotes(query);
if statements.len() <= 1 {
return self.handle_single_query(query).await;
}
self.suppress_ready_for_query = true;
let last_idx = statements.len() - 1;
for (i, stmt) in statements.iter().enumerate() {
if i == last_idx {
self.suppress_ready_for_query = false;
}
self.handle_single_query(stmt).await?;
}
self.suppress_ready_for_query = false;
Ok(())
}
#[allow(clippy::indexing_slicing)]
pub(super) async fn handle_single_query(&mut self, query: &str) -> Result<()> {
tracing::debug!("Executing query: {}", query);
if pg_looks_like_do_block(query.trim()) {
return self.handle_do_block(query).await;
}
let rewritten = pg_rewrite_empty_projection(query);
let query = rewritten.as_str();
if query.trim().is_empty() {
self.send_message(BackendMessage::EmptyQueryResponse).await?;
self.send_ready_for_query().await?;
return Ok(());
}
if let Some(copy_stmt) = super::copy::parse_copy(query) {
return self.handle_copy(copy_stmt).await;
}
let trimmed = query.trim();
if trimmed.eq_ignore_ascii_case("BEGIN")
|| starts_with_icase(trimmed, "BEGIN ")
|| trimmed.eq_ignore_ascii_case("START TRANSACTION")
|| starts_with_icase(trimmed, "START TRANSACTION ")
{
let isolation_level = Self::parse_isolation_level(trimmed);
if self.transaction_status == TransactionStatus::InTransaction {
self.send_message(BackendMessage::NoticeResponse {
severity: "WARNING".to_string(),
code: "25001".to_string(),
message: "there is already a transaction in progress".to_string(),
})
.await?;
} else {
self.rollback_failed_transaction_for_recovery()?;
if let Some(level) = isolation_level.as_deref() {
let mapped = match level {
"SERIALIZABLE" => crate::session::IsolationLevel::Serializable,
"REPEATABLE READ" => crate::session::IsolationLevel::RepeatableRead,
_ => crate::session::IsolationLevel::ReadCommitted,
};
let _ = self.database.set_session_isolation(self.session_id, mapped);
tracing::debug!("Transaction starting with isolation level: {}", level);
}
self.database.begin_transaction_for_session(self.session_id)?;
self.transaction_status = TransactionStatus::InTransaction;
}
self.send_command_complete("BEGIN").await?;
self.send_ready_for_query().await?;
return Ok(());
} else if starts_with_icase(trimmed, "SET TRANSACTION ISOLATION LEVEL ")
|| starts_with_icase(trimmed, "SET SESSION CHARACTERISTICS")
{
let level = Self::parse_isolation_level(trimmed);
if level.is_some() {
self.send_command_complete("SET").await?;
} else {
self.send_error("ERROR", "22023", "Invalid isolation level", None, None)
.await?;
return Ok(());
}
self.send_ready_for_query().await?;
return Ok(());
} else if starts_with_icase(trimmed, "SET ")
&& !starts_with_icase(trimmed, "SET TRANSACTION")
&& !starts_with_icase(trimmed, "SET SESSION CHARACTERISTICS")
{
match EmbeddedDatabase::parse_synchronous_commit_statement(trimmed) {
Ok(Some(value)) => {
if let Err(e) = self.database.set_session_synchronous_commit(self.session_id, value) {
self.send_error("ERROR", "22023", &e.to_string(), None, None).await?;
return Ok(());
}
self.send_command_complete("SET").await?;
self.send_ready_for_query().await?;
return Ok(());
}
Err(e) => {
self.send_error("ERROR", "22023", &e.to_string(), None, None).await?;
return Ok(());
}
Ok(None) => {}
}
match EmbeddedDatabase::parse_helios_fast_autocommit_statement(trimmed) {
Ok(Some(value)) => {
if let Err(e) = self.database.set_session_fast_autocommit(self.session_id, value) {
self.send_error("ERROR", "22023", &e.to_string(), None, None).await?;
return Ok(());
}
self.send_parameter_status("helios.fast_autocommit", if value { "on" } else { "off" })
.await?;
self.send_command_complete("SET").await?;
self.send_ready_for_query().await?;
return Ok(());
}
Err(e) => {
self.send_error("ERROR", "22023", &e.to_string(), None, None).await?;
return Ok(());
}
Ok(None) => {}
}
if EmbeddedDatabase::is_fk_setting_statement(trimmed) {
if let Err(e) = self.database.execute(trimmed) {
self.send_error("ERROR", "22023", &e.to_string(), None, None).await?;
return Ok(());
}
}
self.send_command_complete("SET").await?;
self.send_ready_for_query().await?;
return Ok(());
} else if starts_with_icase(trimmed, "RESET ")
&& trimmed[6..]
.trim()
.trim_end_matches(';')
.trim()
.eq_ignore_ascii_case("synchronous_commit")
{
if let Err(e) = self.database.set_session_synchronous_commit(self.session_id, None) {
self.send_error("ERROR", "22023", &e.to_string(), None, None).await?;
return Ok(());
}
self.send_command_complete("RESET").await?;
self.send_ready_for_query().await?;
return Ok(());
} else if starts_with_icase(trimmed, "RESET ")
&& trimmed[6..]
.trim()
.trim_end_matches(';')
.trim()
.eq_ignore_ascii_case("helios.fast_autocommit")
{
let _ = self.database.set_session_fast_autocommit(self.session_id, false);
self.send_parameter_status("helios.fast_autocommit", "off").await?; self.send_command_complete("RESET").await?;
self.send_ready_for_query().await?;
return Ok(());
} else if starts_with_icase(trimmed, "SHOW ") && !crate::sql::Parser::is_show_branches(trimmed) {
let param = trimmed[5..].trim().trim_end_matches(';').trim();
let (col_name, value) = if param.eq_ignore_ascii_case("synchronous_commit") {
let effective = self
.database
.session_synchronous_commit_effective(self.session_id)
.unwrap_or(false);
(
"synchronous_commit".to_string(),
if effective { "on".to_string() } else { "off".to_string() },
)
} else if param.eq_ignore_ascii_case("helios.fast_autocommit") {
let on = self
.database
.session_fast_autocommit(self.session_id)
.unwrap_or(false);
(
"helios.fast_autocommit".to_string(),
if on { "on".to_string() } else { "off".to_string() },
)
} else {
Self::resolve_show_parameter(param)
};
let schema = Schema::new(vec![crate::Column::new(&col_name, crate::DataType::Text)]);
let row = Tuple::new(vec![Value::String(value)]);
let rows = vec![row];
self.send_query_result(schema, &rows).await?;
self.send_ready_for_query().await?;
return Ok(());
} else if starts_with_icase(trimmed, "PRAGMA ") || trimmed.eq_ignore_ascii_case("PRAGMA") {
if let Some((name, arg)) = crate::sql::sqlite_compat::parse_pragma(trimmed) {
match name.to_lowercase().as_str() {
"table_info" => {
let table = arg.unwrap_or_default();
let table = table
.trim()
.trim_matches(|c| c == '\'' || c == '"' || c == '`')
.to_string();
let rows = self.pragma_table_info(&table)?;
let schema = Schema::new(vec![
crate::Column::new("cid", crate::DataType::Int4),
crate::Column::new("name", crate::DataType::Text),
crate::Column::new("type", crate::DataType::Text),
crate::Column::new("notnull", crate::DataType::Int4),
crate::Column::new("dflt_value", crate::DataType::Text),
crate::Column::new("pk", crate::DataType::Int4),
]);
self.send_query_result(schema, &rows).await?;
self.send_ready_for_query().await?;
return Ok(());
}
_ => {
tracing::debug!("PRAGMA stubbed (no-op): {} = {:?}", name, arg);
self.send_command_complete("PRAGMA").await?;
self.send_ready_for_query().await?;
return Ok(());
}
}
} else {
self.send_command_complete("PRAGMA").await?;
self.send_ready_for_query().await?;
return Ok(());
}
} else if trimmed.eq_ignore_ascii_case("COMMIT") {
if self.transaction_status == TransactionStatus::InTransaction {
self.database.commit_transaction_for_session(self.session_id)?;
} else if self.transaction_status == TransactionStatus::Failed {
self.rollback_failed_transaction_for_recovery()?;
self.send_command_complete("ROLLBACK").await?;
self.send_ready_for_query().await?;
return Ok(());
} else {
self.send_message(BackendMessage::NoticeResponse {
severity: "WARNING".to_string(),
code: "25P01".to_string(),
message: "there is no transaction in progress".to_string(),
})
.await?;
}
self.transaction_status = TransactionStatus::Idle;
self.send_command_complete("COMMIT").await?;
self.send_ready_for_query().await?;
return Ok(());
} else if trimmed.eq_ignore_ascii_case("ROLLBACK") {
if matches!(
self.transaction_status,
TransactionStatus::InTransaction | TransactionStatus::Failed
) {
if self.database.session_in_transaction(self.session_id) {
self.database.rollback_transaction_for_session(self.session_id)?;
}
} else {
self.send_message(BackendMessage::NoticeResponse {
severity: "WARNING".to_string(),
code: "25P01".to_string(),
message: "there is no transaction in progress".to_string(),
})
.await?;
}
self.transaction_status = TransactionStatus::Idle;
self.send_command_complete("ROLLBACK").await?;
self.send_ready_for_query().await?;
return Ok(());
} else if let Some(target) = Self::parse_discard_target(trimmed) {
match target {
"ALL" => {
if matches!(
self.transaction_status,
TransactionStatus::InTransaction | TransactionStatus::Failed
) && self.database.session_in_transaction(self.session_id)
{
self.database.rollback_transaction_for_session(self.session_id)?;
}
self.transaction_status = TransactionStatus::Idle;
self.prepared_statements.clear_all()?;
let _ = self.database.set_session_synchronous_commit(self.session_id, None);
let _ = self.database.set_session_fast_autocommit(self.session_id, false);
let _ = self.database.set_session_isolation(
self.session_id,
crate::session::IsolationLevel::ReadCommitted,
);
self.send_parameter_status("helios.fast_autocommit", "off").await?;
}
"PLANS" => {
self.prepared_statements.clear_all()?;
}
_ => {}
}
self.send_command_complete(&format!("DISCARD {target}")).await?;
self.send_ready_for_query().await?;
return Ok(());
}
if self.transaction_status == TransactionStatus::Failed {
self.send_error(
"ERROR",
"25P02",
"current transaction is aborted, commands ignored until end of transaction block",
None,
Some("Use ROLLBACK to clear the failed transaction state".to_string()),
)
.await?;
return Ok(());
}
#[cfg(feature = "ha-tier1")]
{
use crate::replication::ha_state::{ha_state, SyncMode};
use crate::replication::query_forwarder::{query_forwarder, ForwardedResult};
if ha_state().is_read_only() {
let is_write = starts_with_icase(trimmed, "INSERT")
|| starts_with_icase(trimmed, "UPDATE")
|| starts_with_icase(trimmed, "DELETE")
|| starts_with_icase(trimmed, "CREATE")
|| starts_with_icase(trimmed, "DROP")
|| starts_with_icase(trimmed, "ALTER")
|| starts_with_icase(trimmed, "TRUNCATE");
if is_write {
let config = ha_state().get_config();
let sync_mode = config.as_ref().map(|c| c.sync_mode).unwrap_or(SyncMode::Async);
if matches!(sync_mode, SyncMode::Sync | SyncMode::SemiSync) {
if let Some(forwarder) = query_forwarder() {
match forwarder.forward_query(query) {
Ok(ForwardedResult::Command { tag, .. }) => {
self.send_command_complete(&tag).await?;
self.send_ready_for_query().await?;
return Ok(());
}
Ok(ForwardedResult::Rows { columns, rows }) => {
self.send_forwarded_rows(&columns, &rows).await?;
self.send_ready_for_query().await?;
return Ok(());
}
Ok(ForwardedResult::Error {
severity,
code,
message,
detail,
hint,
}) => {
self.send_error(&severity, &code, &message, detail, hint).await?;
self.send_ready_for_query().await?;
return Ok(());
}
Err(e) => {
self.send_error(
"ERROR",
"08006",
&format!("Failed to forward query to primary: {}", e),
None,
Some("Check primary connectivity".to_string()),
)
.await?;
self.send_ready_for_query().await?;
return Ok(());
}
}
} else {
self.send_error(
"ERROR",
"25006",
"cannot execute write operations: primary connection not established",
None,
Some("Standby is still connecting to primary".to_string()),
)
.await?;
self.send_ready_for_query().await?;
return Ok(());
}
} else {
self.send_error(
"ERROR",
"25006",
"cannot execute write operations in read-only mode (async standby)",
None,
Some("Connect to the primary for write operations, or configure sync mode for transparent routing.".to_string()),
).await?;
self.send_ready_for_query().await?;
return Ok(());
}
}
}
}
if let Some(result) = self.catalog.handle_query(query)? {
let (schema, rows) = result;
self.send_query_result(schema, &rows).await?;
self.send_ready_for_query().await?;
return Ok(());
}
let is_select = starts_with_icase(trimmed, "SELECT");
let is_show_branches = crate::sql::Parser::is_show_branches(trimmed);
let is_cte = !is_select && starts_with_cte(trimmed);
let is_dml_returning = !is_select && !is_cte && {
let upper = trimmed.to_uppercase();
(starts_with_icase(trimmed, "INSERT")
|| starts_with_icase(trimmed, "UPDATE")
|| starts_with_icase(trimmed, "DELETE"))
&& upper.contains("RETURNING")
};
if is_select || is_show_branches {
let cached_query = if is_show_branches || self.database.session_in_transaction(self.session_id) {
None
} else {
self.database.try_cached_query_with_columns(query)
};
if let Some((cached_results, columns)) = cached_query {
let schema = Self::schema_from_query_columns(&columns, cached_results.as_slice());
self.send_query_result(schema, cached_results.as_slice()).await?;
} else {
let (results, columns) = self.database.query_with_columns_for_session(self.session_id, query)?;
let schema = Self::schema_from_query_columns(&columns, &results);
self.send_query_result(schema, &results).await?;
}
} else if is_cte {
let (results, columns) = self.database.query_with_columns_for_session(self.session_id, query)?;
let schema = Self::schema_from_query_columns(&columns, &results);
self.send_query_result(schema, &results).await?;
} else if is_dml_returning {
let (affected, tuples) = self.database.execute_returning_for_session(self.session_id, query)?;
if tuples.is_empty() {
let tag = self.get_command_tag(query, affected);
self.send_command_complete(&tag).await?;
} else {
let schema = self.derive_returning_schema(query).unwrap_or_else(|_| {
if let Some(first) = tuples.first() {
first.schema()
} else {
Schema::new(vec![])
}
});
self.send_query_result(schema, &tuples).await?;
}
} else {
let affected = self.database.execute_for_session(self.session_id, query)?;
let tag = self.get_command_tag(query, affected);
self.send_command_complete(&tag).await?;
}
self.send_ready_for_query().await?;
Ok(())
}
fn schema_from_query_columns(columns: &[String], rows: &[Tuple]) -> Schema {
if !columns.is_empty() {
Schema::new(
columns
.iter()
.enumerate()
.map(|(i, name)| {
let data_type = rows
.first()
.and_then(|r| r.values.get(i))
.map(Value::data_type)
.unwrap_or(crate::DataType::Text);
crate::Column {
name: name.clone(),
data_type,
nullable: true,
primary_key: false,
source_table: None,
source_table_name: None,
default_expr: None,
unique: false,
storage_mode: crate::ColumnStorageMode::Default,
}
})
.collect(),
)
} else if !rows.is_empty() {
rows[0].schema()
} else {
Schema::new(vec![])
}
}
#[allow(clippy::indexing_slicing)]
async fn handle_scram_authentication(&mut self) -> Result<()> {
self.send_message(BackendMessage::Authentication(AuthenticationMessage::ScramSha256))
.await?;
self.flush().await?;
let client_first = match self.read_message().await? {
Some(FrontendMessage::SaslInitialResponse { mechanism: _, data }) => String::from_utf8(data)
.map_err(|e| Error::protocol(format!("SCRAM client-first-message not UTF-8: {}", e)))?,
Some(FrontendMessage::PasswordMessage { password }) => password,
_ => return Err(Error::protocol("Expected SASL initial response")),
};
tracing::debug!("Received client-first-message: {}", client_first);
let (_scram_user_unused, client_nonce_owned) = super::auth::parse_scram_client_first(&client_first)?;
let username_owned = self
.username
.clone()
.ok_or_else(|| Error::authentication("SCRAM authentication started without a startup-time user name"))?;
let username: &str = &username_owned;
let client_nonce: &str = &client_nonce_owned;
tracing::info!("SCRAM authentication for user: {}", username);
let password_store = self
.auth_manager
.password_store()
.ok_or_else(|| Error::authentication("SCRAM password store not configured"))?;
let credentials = password_store
.get_credentials(username)
.ok_or_else(|| Error::authentication("User not found"))?;
let mut scram_state = ScramAuthState::new(username.to_string());
scram_state.set_client_nonce(client_nonce.to_string());
let client_first_bare = format!("n=,r={}", client_nonce);
scram_state.set_client_first_message_bare(client_first_bare);
let server_first = scram_state.build_server_first_message()?;
tracing::debug!("Sending server-first-message: {}", server_first);
self.send_message(BackendMessage::Authentication(
AuthenticationMessage::ScramSha256Continue {
data: server_first.as_bytes().to_vec(),
},
))
.await?;
self.flush().await?;
let client_final = match self.read_message().await? {
Some(FrontendMessage::SaslResponse { data }) => String::from_utf8(data)
.map_err(|e| Error::protocol(format!("SCRAM client-final-message not UTF-8: {}", e)))?,
Some(FrontendMessage::PasswordMessage { password }) => password,
_ => return Err(Error::protocol("Expected SASL response")),
};
tracing::debug!("Received client-final-message: {}", client_final);
let final_parts: Vec<&str> = client_final.split(',').collect();
if final_parts.len() < 3 {
return Err(Error::protocol("Invalid SCRAM client-final-message"));
}
let proof_part = final_parts
.iter()
.find(|p| p.starts_with("p="))
.ok_or_else(|| Error::protocol("Missing proof in client-final-message"))?;
let client_proof_b64 = proof_part
.strip_prefix("p=")
.ok_or_else(|| Error::protocol("Invalid proof format"))?;
let client_final_without_proof: Vec<&str> =
final_parts.iter().filter(|p| !p.starts_with("p=")).copied().collect();
let client_final_without_proof = client_final_without_proof.join(",");
let server_signature = scram_state.verify_client_proof(
client_proof_b64,
&client_final_without_proof,
&credentials.stored_key,
&credentials.server_key,
)?;
tracing::info!("SCRAM authentication successful for user: {}", username);
let server_final = scram_state.build_server_final_message(&server_signature)?;
tracing::debug!("Sending server-final-message: {}", server_final);
self.send_message(BackendMessage::Authentication(
AuthenticationMessage::ScramSha256Final {
data: server_final.as_bytes().to_vec(),
},
))
.await?;
self.authenticated = true;
self.username = Some(username.to_string());
Ok(())
}
async fn send_query_result(&mut self, schema: Schema, rows: &[Tuple]) -> Result<()> {
let fields = schema_to_field_descriptions(&schema);
self.send_message(BackendMessage::RowDescription { fields }).await?;
self.send_data_rows_direct(rows).await?;
let tag = format!("SELECT {}", rows.len());
self.send_command_complete(&tag).await?;
Ok(())
}
pub(super) async fn send_data_rows_with_formats(&mut self, rows: &[Tuple], result_formats: &[i16]) -> Result<()> {
if formats_request_text_only(result_formats) {
return self.send_data_rows_direct(rows).await;
}
for row in rows {
let values = tuple_to_pg_values_with_formats(row, result_formats);
self.send_message(BackendMessage::DataRow { values }).await?;
}
Ok(())
}
pub(super) async fn send_data_rows_direct(&mut self, rows: &[Tuple]) -> Result<()> {
const DATA_ROW_FLUSH_AT: usize = 64 * 1024;
self.write_buf.clear();
for row in rows {
self.encode_data_row_direct(row);
if self.write_buf.len() >= DATA_ROW_FLUSH_AT {
self.stream
.write_all(&self.write_buf)
.await
.map_err(|e| Error::network(format!("Failed to send query rows: {}", e)))?;
self.write_buf.clear();
}
}
if !self.write_buf.is_empty() {
self.stream
.write_all(&self.write_buf)
.await
.map_err(|e| Error::network(format!("Failed to send query rows: {}", e)))?;
self.write_buf.clear();
}
Ok(())
}
#[cfg(feature = "ha-tier1")]
async fn send_forwarded_rows(
&mut self,
columns: &[crate::replication::query_forwarder::ColumnInfo],
rows: &[Vec<Option<String>>],
) -> Result<()> {
use crate::protocol::postgres::messages::FieldDescription;
let fields: Vec<FieldDescription> = columns
.iter()
.map(|col| FieldDescription {
name: col.name.clone(),
table_oid: 0,
column_attr_num: 0,
data_type_oid: col.type_oid,
data_type_size: -1,
type_modifier: -1,
format_code: 0, })
.collect();
self.send_message(BackendMessage::RowDescription { fields }).await?;
for row in rows {
let values: Vec<Option<Vec<u8>>> = row.iter().map(|v| v.as_ref().map(|s| s.as_bytes().to_vec())).collect();
self.send_message(BackendMessage::DataRow { values }).await?;
}
let tag = format!("SELECT {}", rows.len());
self.send_command_complete(&tag).await?;
Ok(())
}
async fn handle_copy(&mut self, copy: super::copy::CopyStatement) -> Result<()> {
use super::copy::CopyFormat;
if copy.format == CopyFormat::Binary {
return self
.send_error("ERROR", "0A000", "COPY binary format is not yet supported", None, None)
.await;
}
if copy.to_stdout {
return self.handle_copy_to_stdout(©).await;
}
let ncols = copy.columns.len();
self.send_message(BackendMessage::CopyInResponse {
overall_format: 0,
column_formats: vec![0i16; ncols],
})
.await?;
self.flush().await?;
let mut data: Vec<u8> = Vec::new();
let mut client_fail: Option<String> = None;
loop {
match self.read_message().await? {
Some(FrontendMessage::CopyData(chunk)) => data.extend_from_slice(&chunk),
Some(FrontendMessage::CopyDone) => break,
Some(FrontendMessage::CopyFail(msg)) => {
client_fail = Some(msg);
break;
}
None | Some(FrontendMessage::Terminate) => return Ok(()),
Some(_) => {
return self
.send_error(
"ERROR",
"08P01",
"unexpected message during COPY FROM STDIN",
None,
None,
)
.await;
}
}
}
if let Some(msg) = client_fail {
return self
.send_error("ERROR", "57014", &format!("COPY from stdin failed: {msg}"), None, None)
.await;
}
let rows = if copy.format == CopyFormat::Csv {
super::copy::parse_csv_rows(&data)
} else {
super::copy::parse_text_rows(&data)
};
let total = rows.len();
const BATCH: usize = 500;
for chunk in rows.chunks(BATCH) {
if let Some(sql) = super::copy::build_insert_sql(©.table, ©.columns, chunk) {
if let Err(e) = self.database.execute(&sql) {
return self
.send_error("ERROR", "XX000", &format!("COPY insert failed: {e}"), None, None)
.await;
}
}
}
self.send_command_complete(&format!("COPY {total}")).await?;
self.send_ready_for_query().await
}
async fn handle_copy_to_stdout(&mut self, copy: &super::copy::CopyStatement) -> Result<()> {
let cols_sql = if copy.columns.is_empty() {
"*".to_string()
} else {
copy.columns
.iter()
.map(|c| format!("\"{}\"", c.replace('"', "\"\"")))
.collect::<Vec<_>>()
.join(", ")
};
let sql = format!("SELECT {} FROM \"{}\"", cols_sql, copy.table.replace('"', "\"\""));
let (rows, columns) = match self
.database
.query_with_columns_for_session(self.session_id, &sql)
{
Ok(r) => r,
Err(e) => {
return self
.send_error("ERROR", "XX000", &format!("COPY TO STDOUT failed: {e}"), None, None)
.await;
}
};
let ncols = columns.len();
self.send_message(BackendMessage::CopyOutResponse {
overall_format: 0,
column_formats: vec![0i16; ncols],
})
.await?;
let total = rows.len();
for row in &rows {
let fields = tuple_to_pg_values(row);
let line = if copy.format == super::copy::CopyFormat::Csv {
super::copy::encode_csv_row(&fields)
} else {
super::copy::encode_text_row(&fields)
};
self.send_message(BackendMessage::CopyData(line)).await?;
}
self.send_message(BackendMessage::CopyDone).await?;
self.send_command_complete(&format!("COPY {total}")).await?;
self.send_ready_for_query().await
}
pub(super) async fn send_message(&mut self, msg: BackendMessage) -> Result<()> {
self.write_buf.clear();
msg.encode(&mut self.write_buf);
self.stream
.write_all(&self.write_buf)
.await
.map_err(|e| Error::network(format!("Failed to send message: {}", e)))?;
Ok(())
}
#[allow(clippy::indexing_slicing)] fn encode_data_row_direct(&mut self, tuple: &Tuple) {
self.write_buf.put_u8(b'D');
let length_pos = self.write_buf.len();
self.write_buf.put_i32(0);
self.write_buf.put_i16(tuple.values.len() as i16);
let mut itoa_buf = itoa::Buffer::new();
let mut ryu_buf = ryu::Buffer::new();
for val in &tuple.values {
match val {
Value::Null => {
self.write_buf.put_i32(-1);
}
Value::Boolean(b) => {
self.write_buf.put_i32(1);
self.write_buf.put_u8(if *b { b't' } else { b'f' });
}
Value::Int2(i) => {
let s = itoa_buf.format(*i);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Int4(i) => {
let s = itoa_buf.format(*i);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Int8(i) => {
let s = itoa_buf.format(*i);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Float4(f) => {
let s = ryu_buf.format(*f);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Float8(f) => {
let s = ryu_buf.format(*f);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::String(s) => {
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Bytes(b) => {
self.write_buf.put_i32(b.len() as i32);
self.write_buf.put_slice(b);
}
Value::Json(j) => {
let s = j.to_string();
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Numeric(n) => {
self.write_buf.put_i32(n.len() as i32);
self.write_buf.put_slice(n.as_bytes());
}
Value::Uuid(u) => {
let s = u.to_string();
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Timestamp(ts) => {
let s = ts.naive_utc().format("%Y-%m-%d %H:%M:%S%.6f").to_string();
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Date(d) => {
let s = d.format("%Y-%m-%d").to_string();
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Time(t) => {
let s = t.format("%H:%M:%S%.6f").to_string();
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Interval(micros) => {
let total_secs = micros / 1_000_000;
let days = total_secs / 86400;
let hours = (total_secs % 86400) / 3600;
let mins = (total_secs % 3600) / 60;
let secs = total_secs % 60;
let s = if days > 0 {
format!("{} days {:02}:{:02}:{:02}", days, hours, mins, secs)
} else {
format!("{:02}:{:02}:{:02}", hours, mins, secs)
};
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::Array(arr) => {
let val_length_pos = self.write_buf.len();
self.write_buf.put_i32(0);
self.write_buf.put_u8(b'{');
for (i, v) in arr.iter().enumerate() {
if i > 0 {
self.write_buf.put_u8(b',');
}
match v {
Value::String(s) => {
self.write_buf.put_u8(b'"');
self.write_buf.put_slice(s.as_bytes());
self.write_buf.put_u8(b'"');
}
Value::Null => self.write_buf.put_slice(b"NULL"),
other => {
let s = other.to_string();
self.write_buf.put_slice(s.as_bytes());
}
}
}
self.write_buf.put_u8(b'}');
let val_len = (self.write_buf.len() - val_length_pos - 4) as i32;
self.write_buf[val_length_pos..val_length_pos + 4].copy_from_slice(&val_len.to_be_bytes());
}
Value::Vector(v) => {
let val_length_pos = self.write_buf.len();
self.write_buf.put_i32(0);
self.write_buf.put_u8(b'{');
for (i, x) in v.iter().enumerate() {
if i > 0 {
self.write_buf.put_u8(b',');
}
let s = ryu_buf.format(*x);
self.write_buf.put_slice(s.as_bytes());
}
self.write_buf.put_u8(b'}');
let val_len = (self.write_buf.len() - val_length_pos - 4) as i32;
self.write_buf[val_length_pos..val_length_pos + 4].copy_from_slice(&val_len.to_be_bytes());
}
Value::DictRef { dict_id } => {
let s = itoa_buf.format(*dict_id);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::CasRef { hash } => {
let s = hex::encode(hash);
self.write_buf.put_i32(s.len() as i32);
self.write_buf.put_slice(s.as_bytes());
}
Value::ColumnarRef => {
self.write_buf.put_i32(10);
self.write_buf.put_slice(b"<columnar>");
}
}
}
let msg_len = (self.write_buf.len() - length_pos) as i32;
self.write_buf[length_pos..length_pos + 4].copy_from_slice(&msg_len.to_be_bytes());
}
async fn flush(&mut self) -> Result<()> {
self.stream
.flush()
.await
.map_err(|e| Error::network(format!("Failed to flush stream: {}", e)))
}
async fn send_auth_ok(&mut self) -> Result<()> {
self.send_message(BackendMessage::Authentication(AuthenticationMessage::Ok))
.await
}
async fn send_parameter_status(&mut self, name: &str, value: &str) -> Result<()> {
self.send_message(BackendMessage::ParameterStatus {
name: name.to_string(),
value: value.to_string(),
})
.await
}
async fn handle_do_block(&mut self, query: &str) -> Result<()> {
let body = pg_extract_do_block_body(query).unwrap_or("");
let stripped = pg_strip_begin_end(body.trim());
let (main_body, exception_codes) = pg_split_exception(stripped);
if let Some(kw) = pg_detect_plpgsql(main_body) {
return Err(Error::query_execution(format!(
"PL/pgSQL control flow (`{kw}`) inside DO blocks is not yet \
supported in HeliosDB Nano. Rewrite the block as plain SQL, \
or execute each statement separately. \
See: docs/compatibility/plpgsql.md"
)));
}
let statements = pg_split_sql_respecting_quotes(main_body);
if statements.is_empty() {
self.send_command_complete("DO").await?;
self.send_ready_for_query().await?;
return Ok(());
}
let prev = self.suppress_ready_for_query;
self.suppress_ready_for_query = true;
for stmt in &statements {
if let Err(e) = self.database.execute_for_session(self.session_id, stmt) {
if pg_exception_matches(&exception_codes, &e.to_string()) {
tracing::debug!("DO block: caught {:?} via EXCEPTION clause; continuing", e.to_string());
continue;
}
self.suppress_ready_for_query = prev;
return Err(e);
}
}
self.suppress_ready_for_query = prev;
self.send_command_complete("DO").await?;
self.send_ready_for_query().await?;
Ok(())
}
async fn send_ready_for_query(&mut self) -> Result<()> {
if self.suppress_ready_for_query {
return Ok(());
}
self.send_message(BackendMessage::ReadyForQuery {
status: self.transaction_status,
})
.await?;
self.flush().await
}
pub(super) async fn send_command_complete(&mut self, tag: &str) -> Result<()> {
self.send_message(BackendMessage::CommandComplete { tag: tag.to_string() })
.await
}
async fn send_error_message(
&mut self,
severity: &str,
code: &str,
message: &str,
detail: Option<String>,
hint: Option<String>,
) -> Result<()> {
self.send_message(BackendMessage::ErrorResponse {
severity: severity.to_string(),
code: code.to_string(),
message: message.to_string(),
detail,
hint,
position: None,
})
.await
}
async fn send_error(
&mut self,
severity: &str,
code: &str,
message: &str,
detail: Option<String>,
hint: Option<String>,
) -> Result<()> {
self.mark_transaction_failed_after_error();
self.send_error_message(severity, code, message, detail, hint).await?;
self.send_ready_for_query().await
}
async fn send_error_for_query(&mut self, error: &Error, wait_for_sync: bool) -> Result<()> {
self.mark_transaction_failed_after_error();
self.suppress_ready_for_query = false;
let code = sqlstate_for_error(error);
let (detail, hint) = detail_hint_for_error(code, error);
if wait_for_sync {
self.send_error_message("ERROR", code, &error.to_string(), detail, hint)
.await?;
self.flush().await
} else {
self.send_error("ERROR", code, &error.to_string(), detail, hint).await
}
}
fn mark_transaction_failed_after_error(&mut self) {
if self.transaction_status == TransactionStatus::InTransaction {
self.transaction_status = TransactionStatus::Failed;
}
}
fn rollback_failed_transaction_for_recovery(&mut self) -> Result<()> {
if self.transaction_status == TransactionStatus::Failed {
if self.database.session_in_transaction(self.session_id) {
self.database.rollback_transaction_for_session(self.session_id)?;
}
self.transaction_status = TransactionStatus::Idle;
}
Ok(())
}
pub(super) fn transaction_failed(&self) -> bool {
self.transaction_status == TransactionStatus::Failed
}
pub(super) async fn send_extended_failed_transaction_error(&mut self) -> Result<()> {
self.suppress_ready_for_query = false;
self.awaiting_sync_after_error = true;
self.send_error_message(
"ERROR",
"25P02",
"current transaction is aborted, commands ignored until end of transaction block",
None,
Some("Use ROLLBACK to clear the failed transaction state".to_string()),
)
.await?;
self.flush().await
}
fn pragma_table_info(&self, table: &str) -> Result<Vec<Tuple>> {
let catalog = self.database.storage.catalog();
let schema = catalog.get_table_schema(table)?;
let mut rows = Vec::with_capacity(schema.columns.len());
for (idx, col) in schema.columns.iter().enumerate() {
rows.push(Tuple::new(vec![
Value::Int4(idx as i32),
Value::String(col.name.clone()),
Value::String(format!("{:?}", col.data_type).to_uppercase()),
Value::Int4(if col.nullable { 0 } else { 1 }),
col.default_expr
.as_ref()
.map(|d| Value::String(d.clone()))
.unwrap_or(Value::Null),
Value::Int4(if col.primary_key { 1 } else { 0 }),
]));
}
Ok(rows)
}
fn derive_returning_schema(&self, sql: &str) -> Result<Schema> {
let catalog = self.database.storage.catalog();
let planner = crate::sql::planner::Planner::with_catalog(&catalog).with_sql(sql.to_string());
let (statement, _) = self.database.parse_cached(sql)?;
let plan = planner.statement_to_plan(statement)?;
let (table_name, returning_items) = match &plan {
crate::sql::LogicalPlan::Insert {
table_name, returning, ..
}
| crate::sql::LogicalPlan::InsertSelect {
table_name, returning, ..
} => (table_name.as_str(), returning.as_ref()),
crate::sql::LogicalPlan::Update {
table_name, returning, ..
} => (table_name.as_str(), returning.as_ref()),
crate::sql::LogicalPlan::Delete {
table_name, returning, ..
} => (table_name.as_str(), returning.as_ref()),
_ => return Err(crate::Error::query_execution("Not a DML statement")),
};
if let Some(items) = returning_items {
let table_schema = catalog.get_table_schema(table_name)?;
Ok(crate::EmbeddedDatabase::returning_schema(&table_schema, items))
} else {
Ok(Schema::new(vec![]))
}
}
pub(super) fn get_command_tag(&self, query: &str, affected: u64) -> String {
let trimmed = query.trim();
if starts_with_icase(trimmed, "INSERT") {
format!("INSERT 0 {}", affected)
} else if starts_with_icase(trimmed, "UPDATE") {
format!("UPDATE {}", affected)
} else if starts_with_icase(trimmed, "DELETE") {
format!("DELETE {}", affected)
} else if starts_with_icase(trimmed, "CREATE TABLE") {
"CREATE TABLE".to_string()
} else if starts_with_icase(trimmed, "DROP TABLE") {
"DROP TABLE".to_string()
} else if starts_with_icase(trimmed, "CREATE INDEX") {
"CREATE INDEX".to_string()
} else {
format!("OK {}", affected)
}
}
#[allow(clippy::indexing_slicing)]
fn resolve_show_parameter(param: &str) -> (String, String) {
let param_lower = param.to_lowercase();
let col = param_lower.clone();
let val = match param_lower.as_str() {
"server_version" => format!("16.0 (HeliosDB Nano {})", env!("CARGO_PKG_VERSION")),
"server_encoding" => "UTF8".to_string(),
"client_encoding" => "UTF8".to_string(),
"standard_conforming_strings" => "on".to_string(),
"transaction_isolation" | "transaction isolation level" => "read committed".to_string(),
"datestyle" => "ISO, MDY".to_string(),
"timezone" | "time zone" => "UTC".to_string(),
"integer_datetimes" => "on".to_string(),
"max_connections" => "100".to_string(),
"lc_collate" => "en_US.UTF-8".to_string(),
"lc_ctype" => "en_US.UTF-8".to_string(),
"search_path" => "\"$user\", public".to_string(),
"default_transaction_isolation" => "read committed".to_string(),
"is_superuser" => "on".to_string(),
_ => String::new(),
};
(col, val)
}
fn parse_discard_target(trimmed: &str) -> Option<&'static str> {
let rest = trimmed.trim_end_matches(';').trim();
if !starts_with_icase(rest, "DISCARD ") {
return None;
}
let arg = rest[8..].trim(); if arg.eq_ignore_ascii_case("ALL") {
Some("ALL")
} else if arg.eq_ignore_ascii_case("PLANS") {
Some("PLANS")
} else if arg.eq_ignore_ascii_case("SEQUENCES") {
Some("SEQUENCES")
} else if arg.eq_ignore_ascii_case("TEMP") || arg.eq_ignore_ascii_case("TEMPORARY") {
Some("TEMP")
} else {
None
}
}
fn parse_isolation_level(query: &str) -> Option<String> {
let query_bytes = query.as_bytes();
let needle = b"ISOLATION LEVEL";
let pos = query_bytes
.windows(needle.len())
.position(|w| w.eq_ignore_ascii_case(needle))?;
let rest = query[pos + needle.len()..].trim();
if starts_with_icase(rest, "READ UNCOMMITTED") {
Some("READ UNCOMMITTED".to_string())
} else if starts_with_icase(rest, "READ COMMITTED") {
Some("READ COMMITTED".to_string())
} else if starts_with_icase(rest, "REPEATABLE READ") {
Some("REPEATABLE READ".to_string())
} else if starts_with_icase(rest, "SERIALIZABLE") {
Some("SERIALIZABLE".to_string())
} else {
None
}
}
}
pub(super) fn schema_to_field_descriptions(schema: &Schema) -> Vec<FieldDescription> {
schema
.columns
.iter()
.map(|col| {
FieldDescription {
name: col.name.clone(),
table_oid: 0,
column_attr_num: 0,
data_type_oid: datatype_to_oid(&col.data_type),
data_type_size: datatype_to_size(&col.data_type),
type_modifier: -1,
format_code: 0, }
})
.collect()
}
pub(super) fn schema_to_field_descriptions_with_formats(
schema: &Schema,
result_formats: &[i16],
) -> Vec<FieldDescription> {
schema
.columns
.iter()
.enumerate()
.map(|(index, col)| FieldDescription {
name: col.name.clone(),
table_oid: 0,
column_attr_num: 0,
data_type_oid: datatype_to_oid(&col.data_type),
data_type_size: datatype_to_size(&col.data_type),
type_modifier: -1,
format_code: effective_result_format(&col.data_type, result_formats, index),
})
.collect()
}
pub(super) fn formats_request_text_only(result_formats: &[i16]) -> bool {
result_formats.iter().all(|f| *f == 0)
}
pub(super) fn requested_result_format(result_formats: &[i16], column_index: usize) -> i16 {
if result_formats.is_empty() {
0
} else if result_formats.len() == 1 {
result_formats[0]
} else {
result_formats.get(column_index).copied().unwrap_or(0)
}
}
fn effective_result_format(data_type: &crate::DataType, result_formats: &[i16], column_index: usize) -> i16 {
let requested = requested_result_format(result_formats, column_index);
if requested == 1 && datatype_has_binary_result(data_type) {
1
} else {
0
}
}
fn datatype_has_binary_result(data_type: &crate::DataType) -> bool {
matches!(
data_type,
crate::DataType::Boolean
| crate::DataType::Int2
| crate::DataType::Int4
| crate::DataType::Int8
| crate::DataType::Float4
| crate::DataType::Float8
| crate::DataType::Bytea
| crate::DataType::Text
| crate::DataType::Varchar(_)
| crate::DataType::Uuid
)
}
pub(super) fn datatype_to_oid(dt: &crate::DataType) -> i32 {
match dt {
crate::DataType::Boolean => 16,
crate::DataType::Int2 => 21,
crate::DataType::Int4 => 23,
crate::DataType::Int8 => 20,
crate::DataType::Float4 => 700,
crate::DataType::Float8 => 701,
crate::DataType::Numeric => 1700,
crate::DataType::Text => 25,
crate::DataType::Varchar(_) => 1043,
crate::DataType::Char(_) => 1042, crate::DataType::Bytea => 17,
crate::DataType::Json => 114,
crate::DataType::Jsonb => 3802,
crate::DataType::Timestamp => 1114,
crate::DataType::Timestamptz => 1184,
crate::DataType::Interval => 1186,
crate::DataType::Date => 1082,
crate::DataType::Time => 1083,
crate::DataType::Uuid => 2950,
crate::DataType::Vector(_) => 1000, _ => 705, }
}
pub(super) fn datatype_to_size(dt: &crate::DataType) -> i16 {
match dt {
crate::DataType::Boolean => 1,
crate::DataType::Int2 => 2,
crate::DataType::Int4 => 4,
crate::DataType::Int8 => 8,
crate::DataType::Float4 => 4,
crate::DataType::Float8 => 8,
crate::DataType::Text => -1, crate::DataType::Varchar(_) => -1,
crate::DataType::Uuid => 16,
_ => -1,
}
}
pub(super) fn tuple_to_pg_values(tuple: &Tuple) -> Vec<Option<Vec<u8>>> {
tuple
.values
.iter()
.map(|val| {
match val {
Value::Null => None,
Value::Boolean(b) => Some(if *b { b"t" } else { b"f" }.to_vec()),
Value::Int2(i) => Some(itoa::Buffer::new().format(*i).as_bytes().to_vec()),
Value::Int4(i) => Some(itoa::Buffer::new().format(*i).as_bytes().to_vec()),
Value::Int8(i) => Some(itoa::Buffer::new().format(*i).as_bytes().to_vec()),
Value::Float4(f) => Some(ryu::Buffer::new().format(*f).as_bytes().to_vec()),
Value::Float8(f) => Some(ryu::Buffer::new().format(*f).as_bytes().to_vec()),
Value::String(s) => Some(s.as_bytes().to_vec()),
Value::Bytes(b) => Some(b.clone()),
Value::Json(j) => Some(j.to_string().into_bytes()),
Value::Numeric(n) => Some(n.as_bytes().to_vec()),
Value::Uuid(u) => Some(u.to_string().into_bytes()),
Value::Timestamp(ts) => Some(ts.naive_utc().format("%Y-%m-%d %H:%M:%S%.6f").to_string().into_bytes()),
Value::Date(d) => Some(d.format("%Y-%m-%d").to_string().into_bytes()),
Value::Time(t) => Some(t.format("%H:%M:%S%.6f").to_string().into_bytes()),
Value::Interval(micros) => {
let total_secs = micros / 1_000_000;
let days = total_secs / 86400;
let hours = (total_secs % 86400) / 3600;
let mins = (total_secs % 3600) / 60;
let secs = total_secs % 60;
let s = if days > 0 {
format!("{} days {:02}:{:02}:{:02}", days, hours, mins, secs)
} else {
format!("{:02}:{:02}:{:02}", hours, mins, secs)
};
Some(s.into_bytes())
}
Value::Array(arr) => {
let mut buf = String::with_capacity(arr.len() * 8 + 2);
buf.push('{');
for (i, v) in arr.iter().enumerate() {
if i > 0 {
buf.push(',');
}
match v {
Value::String(s) => {
buf.push('"');
buf.push_str(s);
buf.push('"');
}
Value::Null => buf.push_str("NULL"),
other => buf.push_str(&other.to_string()),
}
}
buf.push('}');
Some(buf.into_bytes())
}
Value::Vector(v) => {
let mut buf = String::with_capacity(v.len() * 8 + 2);
buf.push('{');
let mut ryu_buf = ryu::Buffer::new();
for (i, x) in v.iter().enumerate() {
if i > 0 {
buf.push(',');
}
buf.push_str(ryu_buf.format(*x));
}
buf.push('}');
Some(buf.into_bytes())
}
Value::DictRef { dict_id } => Some(itoa::Buffer::new().format(*dict_id).as_bytes().to_vec()),
Value::CasRef { hash } => Some(hex::encode(hash).into_bytes()),
Value::ColumnarRef => Some(b"<columnar>".to_vec()),
}
})
.collect()
}
pub(super) fn tuple_to_pg_values_with_formats(tuple: &Tuple, result_formats: &[i16]) -> Vec<Option<Vec<u8>>> {
tuple
.values
.iter()
.enumerate()
.map(|(index, val)| match val {
Value::Null => None,
_ if requested_result_format(result_formats, index) == 1 => {
value_to_pg_binary(val).or_else(|| single_value_to_pg_text(val))
}
_ => single_value_to_pg_text(val),
})
.collect()
}
fn single_value_to_pg_text(value: &Value) -> Option<Vec<u8>> {
let tuple = Tuple::new(vec![value.clone()]);
tuple_to_pg_values(&tuple).into_iter().next().flatten()
}
fn value_to_pg_binary(value: &Value) -> Option<Vec<u8>> {
match value {
Value::Boolean(value) => Some(vec![u8::from(*value)]),
Value::Int2(value) => Some(value.to_be_bytes().to_vec()),
Value::Int4(value) => Some(value.to_be_bytes().to_vec()),
Value::Int8(value) => Some(value.to_be_bytes().to_vec()),
Value::Float4(value) => Some(value.to_be_bytes().to_vec()),
Value::Float8(value) => Some(value.to_be_bytes().to_vec()),
Value::String(value) => Some(value.as_bytes().to_vec()),
Value::Bytes(value) => Some(value.clone()),
Value::Uuid(value) => Some(value.as_bytes().to_vec()),
_ => None,
}
}
fn pg_split_sql_respecting_quotes(sql: &str) -> Vec<String> {
let mut statements = Vec::new();
let mut current = String::new();
let mut in_single_quote = false;
let mut in_dollar: Option<String> = None; let bytes = sql.as_bytes();
let mut i = 0usize;
while i < bytes.len() {
let b = bytes[i];
if let Some(tag) = &in_dollar {
current.push(b as char);
if b == b'$' {
let close = format!("${tag}$");
if sql.get(i..i + close.len()) == Some(close.as_str()) {
for c in close.chars().skip(1) {
current.push(c);
}
i += close.len();
in_dollar = None;
continue;
}
}
i += 1;
continue;
}
if in_single_quote {
current.push(b as char);
if b == b'\'' {
if bytes.get(i + 1) == Some(&b'\'') {
current.push('\'');
i += 2;
continue;
}
in_single_quote = false;
} else if b == b'\\' {
if let Some(&next) = bytes.get(i + 1) {
current.push(next as char);
i += 2;
continue;
}
}
i += 1;
continue;
}
if b == b'$' {
let rest = &sql[i + 1..];
if let Some(end) = rest.find('$') {
let tag = &rest[..end];
if tag.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
current.push('$');
for c in tag.chars() {
current.push(c);
}
current.push('$');
i += 1 + end + 1;
in_dollar = Some(tag.to_string());
continue;
}
}
}
match b {
b'\'' => {
in_single_quote = true;
current.push('\'');
}
b';' => {
let trimmed = current.trim().to_string();
if !trimmed.is_empty() && !pg_stmt_is_only_comment(&trimmed) {
statements.push(trimmed);
}
current.clear();
}
_ => current.push(b as char),
}
i += 1;
}
let trimmed = current.trim().to_string();
if !trimmed.is_empty() && !pg_stmt_is_only_comment(&trimmed) {
statements.push(trimmed);
}
statements
}
fn pg_stmt_is_only_comment(stmt: &str) -> bool {
for line in stmt.lines() {
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
if !trimmed.starts_with("--") {
return false;
}
}
true
}
fn pg_looks_like_do_block(trimmed: &str) -> bool {
let upper = trimmed.trim_start().to_ascii_uppercase();
upper.starts_with("DO $") || upper.starts_with("DO LANGUAGE ")
}
pub(crate) fn pg_rewrite_empty_projection(query: &str) -> String {
let needle_lower = "select from ";
let mut out = String::with_capacity(query.len() + 16);
let mut remaining = query;
loop {
let lower = remaining.to_ascii_lowercase();
match lower.find(needle_lower) {
None => {
out.push_str(remaining);
break;
}
Some(idx) => {
let after = idx + needle_lower.len();
out.push_str(&remaining[..idx]);
out.push_str(&remaining[idx..idx + 6]); out.push_str(" 1 ");
out.push_str(&remaining[idx + 7..after]); remaining = &remaining[after..];
}
}
}
out
}
fn pg_split_exception(body: &str) -> (&str, Vec<String>) {
let upper = body.to_ascii_uppercase();
let offset = upper
.match_indices("EXCEPTION")
.find(|(idx, _)| {
let before_ok = *idx == 0
|| upper
.as_bytes()
.get(*idx - 1)
.copied()
.is_some_and(|b| b.is_ascii_whitespace());
let after_pos = *idx + "EXCEPTION".len();
let after_ok = after_pos == upper.len()
|| upper
.as_bytes()
.get(after_pos)
.copied()
.is_some_and(|b| b.is_ascii_whitespace());
before_ok && after_ok
})
.map(|(idx, _)| idx);
let Some(offset) = offset else {
return (body, Vec::new());
};
let main = &body[..offset];
let exception_block = &body[offset + "EXCEPTION".len()..];
let mut codes = Vec::new();
let upper_eb = exception_block.to_ascii_uppercase();
let mut search_from = 0;
while let Some(rel) = upper_eb[search_from..].find("WHEN") {
let start = search_from + rel + "WHEN".len();
let after = &upper_eb[start..];
let then_pos = after.find("THEN").unwrap_or(after.len());
let conditions = &after[..then_pos];
for name in conditions.split(|c: char| c == ',' || c == '|' || c == 'O' && false) {
for token in name.split_whitespace() {
let cleaned: String = token
.chars()
.filter(|c| c.is_ascii_alphanumeric() || *c == '_')
.collect();
if cleaned.is_empty() || cleaned.eq_ignore_ascii_case("OR") {
continue;
}
codes.push(cleaned.to_lowercase());
}
}
search_from = start + then_pos.min(after.len());
if search_from >= upper_eb.len() {
break;
}
}
(main, codes)
}
fn pg_exception_matches(codes: &[String], error_message: &str) -> bool {
if codes.is_empty() {
return false;
}
let lower = error_message.to_ascii_lowercase();
for code in codes {
let matches = match code.as_str() {
"duplicate_object" | "duplicate_table" | "duplicate_column" | "duplicate_database" | "duplicate_schema"
| "duplicate_function" | "duplicate_alias" => lower.contains("already exists"),
"unique_violation" => lower.contains("unique") || lower.contains("duplicate key"),
"undefined_table" | "undefined_object" | "undefined_column" | "undefined_function" => {
lower.contains("does not exist") || lower.contains("not found")
}
"others" => true,
_ => false,
};
if matches {
return true;
}
}
false
}
fn pg_detect_plpgsql(body: &str) -> Option<&'static str> {
let upper = body.to_ascii_uppercase();
let scrubbed = upper.replace(" IF NOT EXISTS ", " ").replace(" IF EXISTS ", " ");
let padded = format!(" {scrubbed} ");
const KEYWORDS: &[&str] = &[
" DECLARE ",
" IF ",
" LOOP ",
" FOR ",
" WHILE ",
" RAISE ",
" RETURN ",
" PERFORM ",
" EXIT ",
" CONTINUE ",
];
for kw in KEYWORDS {
if padded.contains(kw) {
let trimmed: &str = kw.trim();
return Some(match trimmed {
"DECLARE" => "DECLARE",
"IF" => "IF",
"LOOP" => "LOOP",
"FOR" => "FOR",
"WHILE" => "WHILE",
"RAISE" => "RAISE",
"RETURN" => "RETURN",
"PERFORM" => "PERFORM",
"EXIT" => "EXIT",
"CONTINUE" => "CONTINUE",
_ => "plpgsql",
});
}
}
if body.contains(":=") {
return Some(":=");
}
None
}
fn pg_extract_do_block_body(sql: &str) -> Option<&str> {
let trimmed = sql.trim();
let after_do = trimmed.get(2..)?.trim_start();
let after_lang = if after_do.to_ascii_uppercase().starts_with("LANGUAGE") {
let after = after_do.get("LANGUAGE".len()..)?.trim_start();
let ident_end = after.find(|c: char| !(c.is_ascii_alphanumeric() || c == '_'))?;
after.get(ident_end..)?.trim_start()
} else {
after_do
};
if !after_lang.starts_with('$') {
return None;
}
let rest = after_lang.get(1..)?;
let tag_end = rest.find('$')?;
let tag = rest.get(..tag_end)?;
let closer = format!("${tag}$");
let body_start_abs = sql.len() - rest.len() + tag_end + 1;
let body_search = sql.get(body_start_abs..)?;
let close_rel = body_search.find(&closer)?;
sql.get(body_start_abs..body_start_abs + close_rel)
}
fn pg_strip_begin_end(body: &str) -> &str {
let mut s = body.trim();
if s.to_ascii_uppercase().starts_with("BEGIN") {
s = s.get(5..).map(str::trim_start).unwrap_or(s);
}
let u = s.to_ascii_uppercase();
for suffix in ["END;", "END"] {
if u.ends_with(suffix) {
s = s.get(..s.len() - suffix.len()).map(str::trim_end).unwrap_or(s);
break;
}
}
s
}
#[cfg(test)]
mod plpgsql_detection_tests {
use super::pg_detect_plpgsql;
#[test]
fn create_table_if_not_exists_is_plain_sql() {
assert_eq!(pg_detect_plpgsql("CREATE TABLE IF NOT EXISTS foo (id integer);"), None,);
assert_eq!(pg_detect_plpgsql("DROP TABLE IF EXISTS foo;"), None,);
assert_eq!(pg_detect_plpgsql("create index if not exists ix_foo on foo(id);"), None,);
}
#[test]
fn real_plpgsql_if_still_detected() {
assert_eq!(pg_detect_plpgsql("IF x > 0 THEN PERFORM 1; END IF;"), Some("IF"),);
}
#[test]
fn other_plpgsql_keywords_still_detected() {
assert_eq!(pg_detect_plpgsql("LOOP exit; END LOOP;"), Some("LOOP"));
assert_eq!(pg_detect_plpgsql("RAISE NOTICE 'hi';"), Some("RAISE"));
assert_eq!(pg_detect_plpgsql("DECLARE v integer;"), Some("DECLARE"));
}
}
#[cfg(test)]
mod datatype_oid_tests {
use super::datatype_to_oid;
use crate::DataType;
#[test]
fn numeric_advertises_numeric_oid_not_unknown() {
assert_eq!(datatype_to_oid(&DataType::Numeric), 1700);
assert_ne!(datatype_to_oid(&DataType::Numeric), 705);
}
#[test]
fn scalar_types_map_to_canonical_oids() {
assert_eq!(datatype_to_oid(&DataType::Char(10)), 1042); assert_eq!(datatype_to_oid(&DataType::Timestamptz), 1184);
assert_eq!(datatype_to_oid(&DataType::Interval), 1186);
assert_eq!(datatype_to_oid(&DataType::Int4), 23);
assert_eq!(datatype_to_oid(&DataType::Int8), 20);
assert_eq!(datatype_to_oid(&DataType::Float8), 701);
assert_eq!(datatype_to_oid(&DataType::Timestamp), 1114);
}
}
pub(crate) fn sqlstate_for_error(error: &Error) -> &'static str {
use crate::network::protocol::sqlstate;
match error {
Error::ConstraintViolation(message) => {
let lower = message.to_ascii_lowercase();
if lower.contains("duplicate") || lower.contains("unique") {
sqlstate::UNIQUE_VIOLATION } else if lower.contains("foreign key") {
sqlstate::FOREIGN_KEY_VIOLATION } else if lower.contains("check") {
sqlstate::CHECK_VIOLATION } else {
sqlstate::INTEGRITY_CONSTRAINT_VIOLATION }
}
Error::SqlParse(_) => sqlstate::SYNTAX_ERROR, Error::TypeConversion(_) => sqlstate::DATATYPE_MISMATCH, Error::Transaction(message) => {
let lower = message.to_ascii_lowercase();
if lower.contains("serialization failure") {
sqlstate::SERIALIZATION_FAILURE } else if lower.contains("deadlock") {
sqlstate::DEADLOCK_DETECTED } else {
sqlstate::INVALID_TRANSACTION_STATE }
}
Error::Protocol(_) => sqlstate::PROTOCOL_VIOLATION, Error::QueryTimeout(_) | Error::QueryCancelled(_) => sqlstate::QUERY_CANCELED, Error::QueryExecution(message) => sqlstate_for_query_execution_message(message),
_ => sqlstate::INTERNAL_ERROR, }
}
fn sqlstate_for_query_execution_message(message: &str) -> &'static str {
use crate::network::protocol::sqlstate;
let lower = message.to_ascii_lowercase();
let not_found = lower.contains("not found") || lower.contains("does not exist") || lower.contains("doesn't exist");
if lower.contains("function") && (not_found || lower.contains("unknown")) {
sqlstate::UNDEFINED_FUNCTION } else if lower.contains("column") && (not_found || lower.contains("unknown")) {
sqlstate::UNDEFINED_COLUMN } else if (lower.contains("table") || lower.contains("relation")) && lower.contains("already exists") {
sqlstate::DUPLICATE_TABLE } else if (lower.contains("table") || lower.contains("relation")) && not_found {
sqlstate::UNDEFINED_TABLE } else {
sqlstate::INTERNAL_ERROR }
}
pub(crate) fn detail_hint_for_error(code: &str, error: &Error) -> (Option<String>, Option<String>) {
use crate::network::protocol::sqlstate;
let message = error.to_string();
match code {
sqlstate::SERIALIZATION_FAILURE => (
Some("Another transaction committed a conflicting write first (first-committer-wins).".to_string()),
Some("Retry the transaction.".to_string()),
),
sqlstate::DEADLOCK_DETECTED => (
Some("This transaction was chosen as the deadlock victim and rolled back.".to_string()),
Some("Retry the transaction.".to_string()),
),
sqlstate::UNDEFINED_TABLE => (
first_single_quoted(&message).map(|name| format!("Table '{name}' does not exist in the catalog.")),
Some("Check the table name, or create the table first.".to_string()),
),
sqlstate::UNDEFINED_COLUMN => (
first_single_quoted(&message).map(|name| format!("Column '{name}' does not exist.")),
Some("Check the column name against the table definition.".to_string()),
),
_ => (None, None),
}
}
fn first_single_quoted(message: &str) -> Option<&str> {
let start = message.find('\'')? + 1;
let rest = message.get(start..)?;
let end = rest.find('\'')?;
rest.get(..end)
}
#[cfg(test)]
mod do_block_split_tests {
use super::pg_split_exception;
#[test]
fn split_handles_newline_before_exception() {
let body = "CREATE TABLE IF NOT EXISTS hdb_test_do (id integer);\n\
EXCEPTION WHEN duplicate_object THEN null;";
let (main, codes) = pg_split_exception(body);
assert!(main.starts_with("CREATE TABLE"));
assert!(!main.to_ascii_uppercase().contains("EXCEPTION"));
assert_eq!(codes, vec!["duplicate_object".to_string()]);
}
#[test]
fn split_still_handles_inline_exception() {
let body = "ALTER TABLE x ADD COLUMN y INT; EXCEPTION WHEN duplicate_column THEN null;";
let (main, codes) = pg_split_exception(body);
assert!(main.starts_with("ALTER TABLE"));
assert!(!main.to_ascii_uppercase().contains("EXCEPTION"));
assert_eq!(codes, vec!["duplicate_column".to_string()]);
}
#[test]
fn split_handles_no_exception() {
let body = "ALTER TABLE x ADD COLUMN y INT;";
let (main, codes) = pg_split_exception(body);
assert_eq!(main, body);
assert!(codes.is_empty());
}
#[test]
fn split_does_not_match_exception_inside_identifier() {
let body = "SELECT * FROM my_exception_table;\nEXCEPTION WHEN others THEN null;";
let (main, codes) = pg_split_exception(body);
assert!(main.starts_with("SELECT"));
assert!(main.to_ascii_uppercase().contains("MY_EXCEPTION_TABLE"));
assert!(!main.to_ascii_uppercase().ends_with("EXCEPTION"));
assert_eq!(codes, vec!["others".to_string()]);
}
}
#[cfg(test)]
mod failed_transaction_state_tests {
use super::*;
use std::sync::Arc;
use tokio::io::DuplexStream;
fn test_handler(db: Arc<EmbeddedDatabase>) -> (PgConnectionHandler<DuplexStream>, DuplexStream) {
let (stream, client) = tokio::io::duplex(4096);
(
PgConnectionHandler {
stream,
session_id: db
.create_wire_session("pg_wire_test")
.expect("wire session creation is infallible"),
database: db.clone(),
auth_manager: Arc::new(AuthManager::new(AuthMethod::Trust)),
catalog: PgCatalog::with_database(db),
prepared_statements: PreparedStatementManager::new(),
authenticated: true,
transaction_status: TransactionStatus::Idle,
buffer: BytesMut::with_capacity(8192),
username: None,
scram_state: None,
write_buf: BytesMut::with_capacity(4096),
suppress_ready_for_query: false,
awaiting_sync_after_error: false,
},
client,
)
}
#[tokio::test]
async fn query_error_marks_open_transaction_failed_and_reenables_ready() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.begin().expect("begin");
let (mut handler, _client) = test_handler(Arc::clone(&db));
handler.transaction_status = TransactionStatus::InTransaction;
handler.suppress_ready_for_query = true;
handler
.send_error_for_query(&Error::constraint_violation("duplicate key"), false)
.await
.expect("send error");
assert_eq!(handler.transaction_status, TransactionStatus::Failed);
assert!(!handler.suppress_ready_for_query);
assert!(db.in_transaction());
handler
.handle_single_query("BEGIN")
.await
.expect("BEGIN recovers failed transaction");
assert_eq!(handler.transaction_status, TransactionStatus::InTransaction);
db.rollback().expect("rollback cleanup");
}
#[tokio::test]
async fn direct_error_response_marks_open_transaction_failed() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.begin().expect("begin");
let (mut handler, _client) = test_handler(Arc::clone(&db));
handler.transaction_status = TransactionStatus::InTransaction;
handler
.send_error("ERROR", "22023", "bad setting", None, None)
.await
.expect("send error");
assert_eq!(handler.transaction_status, TransactionStatus::Failed);
assert!(db.in_transaction());
db.rollback().expect("rollback cleanup");
}
#[tokio::test]
async fn rollback_clears_failed_transaction_state() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("BEGIN").expect("begin");
let (mut handler, _client) = test_handler(db);
handler.transaction_status = TransactionStatus::Failed;
handler
.handle_single_query("ROLLBACK")
.await
.expect("rollback after failed transaction");
assert_eq!(handler.transaction_status, TransactionStatus::Idle);
}
}
#[cfg(test)]
mod show_branches_wire_tests {
use super::*;
use std::sync::Arc;
use std::time::Duration;
use tokio::io::AsyncReadExt;
use tokio::io::DuplexStream;
fn test_handler(db: Arc<EmbeddedDatabase>) -> (PgConnectionHandler<DuplexStream>, DuplexStream) {
let (stream, client) = tokio::io::duplex(4096);
(
PgConnectionHandler {
stream,
session_id: db
.create_wire_session("pg_wire_test")
.expect("wire session creation is infallible"),
database: db.clone(),
auth_manager: Arc::new(AuthManager::new(AuthMethod::Trust)),
catalog: PgCatalog::with_database(db),
prepared_statements: PreparedStatementManager::new(),
authenticated: true,
transaction_status: TransactionStatus::Idle,
buffer: BytesMut::with_capacity(8192),
username: None,
scram_state: None,
write_buf: BytesMut::with_capacity(4096),
suppress_ready_for_query: false,
awaiting_sync_after_error: false,
},
client,
)
}
fn has_complete_ready_for_query(buf: &[u8]) -> bool {
let mut pos = 0;
while pos + 5 <= buf.len() {
let tag = buf[pos];
let len = i32::from_be_bytes([buf[pos + 1], buf[pos + 2], buf[pos + 3], buf[pos + 4]]) as usize;
let end = pos + 1 + len;
if end > buf.len() {
return false;
}
if tag == b'Z' {
return true;
}
pos = end;
}
false
}
async fn read_until_ready(mut client: DuplexStream) -> Vec<u8> {
let mut out = Vec::new();
let mut chunk = [0_u8; 2048];
loop {
let read = tokio::time::timeout(Duration::from_secs(1), client.read(&mut chunk))
.await
.expect("timed out waiting for PostgreSQL response")
.expect("read PostgreSQL response");
assert!(read > 0, "connection closed before ReadyForQuery; bytes={out:?}");
out.extend_from_slice(&chunk[..read]);
if has_complete_ready_for_query(&out) {
return out;
}
}
}
fn first_column_text_values(buf: &[u8]) -> Vec<Option<String>> {
let mut pos = 0;
let mut values = Vec::new();
while pos + 5 <= buf.len() {
let tag = buf[pos];
let len = i32::from_be_bytes([buf[pos + 1], buf[pos + 2], buf[pos + 3], buf[pos + 4]]) as usize;
let end = pos + 1 + len;
if end > buf.len() {
break;
}
if tag == b'D' {
let mut cur = pos + 5;
let columns = i16::from_be_bytes([buf[cur], buf[cur + 1]]) as usize;
cur += 2;
for idx in 0..columns {
let value_len = i32::from_be_bytes([buf[cur], buf[cur + 1], buf[cur + 2], buf[cur + 3]]);
cur += 4;
if value_len < 0 {
if idx == 0 {
values.push(None);
}
continue;
}
let value_end = cur + value_len as usize;
if idx == 0 {
values.push(Some(
std::str::from_utf8(&buf[cur..value_end])
.expect("first column should be UTF-8 text")
.to_string(),
));
}
cur = value_end;
}
}
pos = end;
}
values
}
#[tokio::test]
async fn show_branches_simple_query_returns_branch_registry_rows() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE t(x INT)").expect("create table");
db.execute("INSERT INTO t VALUES(1),(2),(3)").expect("insert rows");
db.execute("CREATE BRANCH 'alpha' AS OF NOW").expect("create alpha");
db.execute("CREATE BRANCH 'beta' AS OF NOW").expect("create beta");
let (mut handler, client) = test_handler(db);
handler
.handle_single_query("SHOW BRANCHES")
.await
.expect("SHOW BRANCHES over PostgreSQL wire");
let response = read_until_ready(client).await;
let names = first_column_text_values(&response);
let rendered: Vec<String> = names
.iter()
.map(|name| name.clone().unwrap_or_else(|| "<NULL>".to_string()))
.collect();
assert!(
rendered.iter().any(|name| name == "main"),
"missing main; names={rendered:?}"
);
assert!(
rendered.iter().any(|name| name == "alpha"),
"missing alpha; names={rendered:?}"
);
assert!(
rendered.iter().any(|name| name == "beta"),
"missing beta; names={rendered:?}"
);
assert!(
rendered.iter().all(|name| !name.is_empty()),
"SHOW BRANCHES must not return a blank branch name; names={rendered:?}"
);
}
}
#[cfg(test)]
mod sqlstate_mapping_unit_tests {
use super::{detail_hint_for_error, first_single_quoted, sqlstate_for_error};
use crate::session::IsolationLevel;
use crate::{EmbeddedDatabase, Error};
fn sql_error(db: &EmbeddedDatabase, sql: &str) -> Error {
db.execute(sql).expect_err("statement must fail")
}
fn query_error(db: &EmbeddedDatabase, sql: &str) -> Error {
db.query(sql, &[]).expect_err("query must fail")
}
#[test]
fn undefined_table_maps_to_42p01_with_detail() {
let db = EmbeddedDatabase::new_in_memory().unwrap();
let err = query_error(&db, "SELECT * FROM no_such_table");
let code = sqlstate_for_error(&err);
assert_eq!(code, "42P01", "got error: {err}");
let (detail, hint) = detail_hint_for_error(code, &err);
assert!(
detail.as_deref().unwrap_or_default().contains("no_such_table"),
"detail must carry the table name; got {detail:?} for {err}"
);
assert!(hint.is_some());
}
#[test]
fn undefined_column_maps_to_42703_with_detail() {
let db = EmbeddedDatabase::new_in_memory().unwrap();
db.execute("CREATE TABLE t42703 (id INTEGER PRIMARY KEY)").unwrap();
db.execute("INSERT INTO t42703 VALUES (1)").unwrap();
let err = query_error(&db, "SELECT no_such_col FROM t42703");
let code = sqlstate_for_error(&err);
assert_eq!(code, "42703", "got error: {err}");
let (detail, hint) = detail_hint_for_error(code, &err);
assert!(
detail.as_deref().unwrap_or_default().contains("no_such_col"),
"detail must carry the column name; got {detail:?} for {err}"
);
assert!(hint.is_some());
}
#[test]
fn duplicate_table_maps_to_42p07() {
let db = EmbeddedDatabase::new_in_memory().unwrap();
db.execute("CREATE TABLE t42p07 (id INTEGER PRIMARY KEY)").unwrap();
let err = sql_error(&db, "CREATE TABLE t42p07 (id INTEGER PRIMARY KEY)");
assert_eq!(sqlstate_for_error(&err), "42P07", "got error: {err}");
}
#[test]
fn undefined_function_maps_to_42883() {
let db = EmbeddedDatabase::new_in_memory().unwrap();
db.execute("CREATE TABLE t42883 (id INTEGER PRIMARY KEY)").unwrap();
db.execute("INSERT INTO t42883 VALUES (1)").unwrap();
let err = query_error(&db, "SELECT definitely_not_a_function(id) FROM t42883");
assert_eq!(sqlstate_for_error(&err), "42883", "got error: {err}");
}
#[test]
fn unique_violation_still_maps_to_23505() {
let db = EmbeddedDatabase::new_in_memory().unwrap();
db.execute("CREATE TABLE t23505 (id INTEGER PRIMARY KEY)").unwrap();
db.execute("INSERT INTO t23505 VALUES (1)").unwrap();
let err = sql_error(&db, "INSERT INTO t23505 VALUES (1)");
assert_eq!(sqlstate_for_error(&err), "23505", "got error: {err}");
}
#[test]
fn serialization_failure_maps_to_40001_with_retry_hint() {
let db = EmbeddedDatabase::new_in_memory().unwrap();
db.execute("CREATE TABLE t40001 (id INTEGER PRIMARY KEY, v INTEGER)")
.unwrap();
db.execute("INSERT INTO t40001 VALUES (1, 100)").unwrap();
let a = db.create_session("a", IsolationLevel::RepeatableRead).unwrap();
let b = db.create_session("b", IsolationLevel::RepeatableRead).unwrap();
db.begin_transaction_for_session(a).unwrap();
db.begin_transaction_for_session(b).unwrap();
db.execute_in_session(a, "UPDATE t40001 SET v = 101 WHERE id = 1")
.unwrap();
db.commit_transaction_for_session(a).unwrap();
db.execute_in_session(b, "UPDATE t40001 SET v = 150 WHERE id = 1")
.unwrap();
let err = db
.commit_transaction_for_session(b)
.expect_err("conflicting commit must fail");
let code = sqlstate_for_error(&err);
assert_eq!(code, "40001", "got error: {err}");
let (detail, hint) = detail_hint_for_error(code, &err);
assert!(detail.is_some());
assert!(
hint.as_deref()
.unwrap_or_default()
.to_ascii_lowercase()
.contains("retry"),
"hint must suggest retrying; got {hint:?}"
);
db.destroy_session(a).unwrap();
db.destroy_session(b).unwrap();
}
#[test]
fn deadlock_maps_to_40p01_with_retry_hint() {
let err = Error::deadlock("Deadlock detected for transaction 7");
let code = sqlstate_for_error(&err);
assert_eq!(code, "40P01");
let (_, hint) = detail_hint_for_error(code, &err);
assert!(hint
.as_deref()
.unwrap_or_default()
.to_ascii_lowercase()
.contains("retry"));
}
#[test]
fn non_serialization_transaction_error_keeps_25000() {
let err = Error::transaction("no transaction in progress");
assert_eq!(sqlstate_for_error(&err), "25000");
}
#[test]
fn unknown_query_execution_error_stays_internal() {
let err = Error::query_execution("something exotic went wrong");
assert_eq!(sqlstate_for_error(&err), "XX000");
}
#[test]
fn quoted_name_extraction() {
assert_eq!(first_single_quoted("Table 'users' does not exist"), Some("users"));
assert_eq!(first_single_quoted("no quotes here"), None);
assert_eq!(first_single_quoted("dangling 'quote"), None);
}
}