diff --git a/CHANGELOG.md b/CHANGELOG.md index 0a1db5f..86997e8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -23,8 +23,40 @@ Versioning](https://semver.org/spec/v2.0.0.html). by PostgreSQL 18+ and the minor-only form used by older servers) and adopts the negotiated version for the rest of the connection. -### Fixed +### Changed +- Breaking: `QueryParser::parse_sql` now returns + `PgWireResult>`, and `StoredStatement::parse` + correspondingly returns `Option>`. `None` denotes an + empty query: it is stored as an empty statement and executes to + `EmptyQueryResponse`. This allows a query parser to report its own notion + of empty query; syntactically empty queries (semicolons and whitespace + only) are still never passed to the parser. +- Breaking: `PortalStore` now represents empty statements and portals. + `get_statement` returns `Option>>` and + `get_portal` returns `Option>>`: the new `Entry::Empty` + variant marks a name under which an empty prepared statement or portal is + stored, alongside the new `put_empty_statement`/`put_empty_portal` + methods. Like every `put_*`, storing an empty entry replaces whatever was + previously stored under that name, and `rm_*`/`clear_portals` remove + empty entries along with regular ones. `StoredStatement` and `Portal` + themselves are unchanged — the impact is limited to `PortalStore` + implementors and code calling `get_statement`/`get_portal` directly + (`Entry::value` helps with the migration). + +### Fixed + +- Extended query protocol: empty queries (a query string without any + statement, such as `""` or `";;"`) are now handled like PostgreSQL instead + of being dispatched to the query parser: `Parse` succeeds without calling + `QueryParser` and stores an empty statement, `Describe` answers + `ParameterDescription` (no parameters) + `NoData`, `Bind` succeeds + (rejecting bound parameters with `08P01`) and stores an empty portal, and + `Execute` returns `EmptyQueryResponse` without reaching `do_query`. An + empty `Parse` replaces any statement previously stored under the same + name, and `Close`/`Sync` drop empty statements and the unnamed empty + portal like real ones. Behavior verified message-for-message against + PostgreSQL 18. - Client API: backend messages are now decoded with the rules of the protocol version the client actually advertised, instead of always 3.2. Previously a 4-byte protocol 3.0 cancel key was decoded as `SecretKey::Bytes` instead of diff --git a/examples/cursor.rs b/examples/cursor.rs index 3f0332e..1b5bca9 100644 --- a/examples/cursor.rs +++ b/examples/cursor.rs @@ -11,7 +11,7 @@ use pgwire::api::portal::Portal; use pgwire::api::query::SimpleQueryHandler; use pgwire::api::results::{DataRowEncoder, FieldFormat, FieldInfo, QueryResponse, Response, Tag}; use pgwire::api::stmt::StoredStatement; -use pgwire::api::store::{MemPortalStore, PortalStore}; +use pgwire::api::store::{Entry, MemPortalStore, PortalStore}; use pgwire::api::{ClientInfo, ClientPortalStore, PgWireServerHandlers, Type}; use pgwire::error::{ErrorInfo, PgWireError, PgWireResult}; use pgwire::messages::response::NoticeResponse; @@ -210,7 +210,7 @@ async fn handle_fetch( ) -> PgWireResult> { println!("FETCH {} FROM {}", count, cursor_name); - let Some(portal) = portal_store.get_portal(cursor_name) else { + let Some(Entry::Value(portal)) = portal_store.get_portal(cursor_name) else { return Err(PgWireError::UserError(Box::new(ErrorInfo::new( "ERROR".to_owned(), "34000".to_owned(), diff --git a/src/api/query.rs b/src/api/query.rs index 67583b4..e0ad942 100644 --- a/src/api/query.rs +++ b/src/api/query.rs @@ -12,7 +12,7 @@ use futures::stream::StreamExt; use super::portal::Portal; use super::results::{Tag, into_row_description}; use super::stmt::{NoopQueryParser, QueryParser, StoredStatement}; -use super::store::PortalStore; +use super::store::{Entry, PortalStore}; use super::{ClientInfo, ClientPortalStore, ConnectionHandle, DEFAULT_NAME, copy}; use crate::api::PgWireConnectionState; use crate::api::Type; @@ -30,7 +30,7 @@ use crate::messages::extendedquery::{ use crate::messages::response::{EmptyQueryResponse, ReadyForQuery, TransactionStatus}; use crate::messages::simplequery::Query; -fn is_empty_query(q: &str) -> bool { +pub(crate) fn is_empty_query(q: &str) -> bool { // A query string that contains only semicolons and whitespace parses to no // statements, which PostgreSQL treats as an empty query and answers with // `EmptyQueryResponse` instead of dispatching to the executor. This covers @@ -186,8 +186,10 @@ pub trait ExtendedQueryHandler: Send + Sync { /// Called when client sends `parse` command. /// - /// The default implementation parsed query with `Self::QueryParser` and - /// stores it in `Self::PortalStore`. + /// The default implementation parses the query with + /// `Self::QueryParser` and stores it in `Self::PortalStore`. Empty + /// queries are stored as empty statements instead, like PostgreSQL: + /// they bind, describe and execute as empty queries. async fn on_parse(&self, client: &mut C, message: Parse) -> PgWireResult<()> where C: ClientInfo + ClientPortalStore + Sink + Unpin + Send + Sync, @@ -195,9 +197,16 @@ pub trait ExtendedQueryHandler: Send + Sync { C::Error: Debug, PgWireError: From<>::Error>, { + let name = message + .name + .clone() + .unwrap_or_else(|| DEFAULT_NAME.to_owned()); + let parser = self.query_parser(); - let stmt = StoredStatement::parse(client, &message, parser).await?; - client.portal_store().put_statement(Arc::new(stmt)); + match StoredStatement::parse(client, &message, parser).await? { + Some(stmt) => client.portal_store().put_statement(Arc::new(stmt)), + None => client.portal_store().put_empty_statement(&name), + } client .send(PgWireBackendMessage::ParseComplete(ParseComplete::new())) .await?; @@ -207,8 +216,10 @@ pub trait ExtendedQueryHandler: Send + Sync { /// Called when client sends `bind` command. /// - /// The default implementation associate parameters with previous parsed - /// statement and stores in `Self::PortalStore` as well. + /// The default implementation associates parameters with a previously + /// parsed statement and stores the result in `Self::PortalStore` as well. + /// Binding an empty statement stores an empty portal, with zero + /// parameters, like PostgreSQL. async fn on_bind(&self, client: &mut C, message: Bind) -> PgWireResult<()> where C: ClientInfo + ClientPortalStore + Sink + Unpin + Send + Sync, @@ -217,27 +228,41 @@ pub trait ExtendedQueryHandler: Send + Sync { PgWireError: From<>::Error>, { let statement_name = message.statement_name.as_deref().unwrap_or(DEFAULT_NAME); + let portal_name = message.portal_name.as_deref().unwrap_or(DEFAULT_NAME); - if let Some(statement) = client.portal_store().get_statement(statement_name) { - let portal = Portal::try_new(&message, statement)?; - client.portal_store().put_portal(Arc::new(portal)); - client - .send(PgWireBackendMessage::BindComplete(BindComplete::new())) - .await?; - Ok(()) - } else { - Err(PgWireError::StatementNotFound(statement_name.to_owned())) + match client.portal_store().get_statement(statement_name) { + Some(Entry::Value(statement)) => { + let portal = Portal::try_new(&message, statement)?; + client.portal_store().put_portal(Arc::new(portal)); + } + Some(Entry::Empty) => { + if !message.parameters.is_empty() { + return Err(PgWireError::UserError(Box::new(ErrorInfo::new( + "ERROR".to_owned(), + "08P01".to_owned(), + format!( + "bind message supplies {} parameters, but prepared statement {:?} requires 0", + message.parameters.len(), + statement_name + ), + )))); + } + client.portal_store().put_empty_portal(portal_name); + } + None => return Err(PgWireError::StatementNotFound(statement_name.to_owned())), } + + client + .send(PgWireBackendMessage::BindComplete(BindComplete::new())) + .await?; + Ok(()) } /// Called when client sends `execute` command. /// /// The default implementation delegates the query to `self::do_query` and /// sends response messages according to `Response` from `self::do_query`. - /// - /// Note that, different from `SimpleQueryHandler`, this implementation - /// won't check empty query because it cannot understand parsed - /// `Self::Statement`. + /// Empty portals answer `EmptyQueryResponse` and never reach `do_query`. async fn on_execute(&self, client: &mut C, message: Execute) -> PgWireResult<()> where C: ClientInfo + ClientPortalStore + Sink + Unpin + Send + Sync, @@ -270,8 +295,17 @@ pub trait ExtendedQueryHandler: Send + Sync { let portal_name = message.name.as_deref().unwrap_or(DEFAULT_NAME); let max_rows = message.max_rows as usize; - let Some(portal) = client.portal_store().get_portal(portal_name) else { - return Err(PgWireError::PortalNotFound(portal_name.to_owned())); + let portal = match client.portal_store().get_portal(portal_name) { + Some(Entry::Value(portal)) => portal, + Some(Entry::Empty) => { + // never reaches do_query; stays valid for repeated Execute + client + .feed(PgWireBackendMessage::EmptyQueryResponse(EmptyQueryResponse)) + .await?; + client.set_state(super::PgWireConnectionState::ReadyForQuery); + return Ok(()); + } + None => return Err(PgWireError::PortalNotFound(portal_name.to_owned())), }; // Execute query if the portal hasn't been started yet let needs_fetch = if matches!( @@ -400,22 +434,28 @@ pub trait ExtendedQueryHandler: Send + Sync { { let name = message.name.as_deref().unwrap_or(DEFAULT_NAME); match message.target_type { - TARGET_TYPE_BYTE_STATEMENT => { - if let Some(stmt) = client.portal_store().get_statement(name) { + TARGET_TYPE_BYTE_STATEMENT => match client.portal_store().get_statement(name) { + Some(Entry::Value(stmt)) => { let describe_response = self.do_describe_statement(client, &stmt).await?; send_describe_response(client, &describe_response).await?; - } else { - return Err(PgWireError::StatementNotFound(name.to_owned())); } - } - TARGET_TYPE_BYTE_PORTAL => { - if let Some(portal) = client.portal_store().get_portal(name) { + Some(Entry::Empty) => { + let describe_response = DescribeStatementResponse::no_data(); + send_describe_response(client, &describe_response).await?; + } + None => return Err(PgWireError::StatementNotFound(name.to_owned())), + }, + TARGET_TYPE_BYTE_PORTAL => match client.portal_store().get_portal(name) { + Some(Entry::Value(portal)) => { let describe_response = self.do_describe_portal(client, &portal).await?; send_describe_response(client, &describe_response).await?; - } else { - return Err(PgWireError::PortalNotFound(name.to_owned())); } - } + Some(Entry::Empty) => { + let describe_response = DescribePortalResponse::no_data(); + send_describe_response(client, &describe_response).await?; + } + None => return Err(PgWireError::PortalNotFound(name.to_owned())), + }, _ => return Err(PgWireError::InvalidTargetType(message.target_type)), } @@ -820,3 +860,666 @@ mod tests { assert!(!is_empty_query("';'")); } } + +/// Extended-query empty statement handling, mirroring PostgreSQL 18 +/// message sequences. +#[cfg(test)] +mod extended_empty_query_tests { + use std::net::SocketAddr; + use std::pin::Pin; + use std::sync::Mutex; + use std::task::{Context, Poll}; + + use async_trait::async_trait; + use bytes::Bytes; + use futures::Sink; + + use super::*; + use crate::api::results::Tag; + use crate::api::{DefaultClient, PgWireConnectionState}; + use crate::messages::response::TransactionStatus; + + /// A client test-double recording backend messages instead of encoding + /// them. + struct TestClient { + inner: DefaultClient, + sent: Mutex>, + } + + impl TestClient { + fn new() -> Self { + let mut inner = DefaultClient::new(SocketAddr::from(([127, 0, 0, 1], 5432)), false); + inner.set_state(PgWireConnectionState::ReadyForQuery); + TestClient { + inner, + sent: Mutex::new(Vec::new()), + } + } + + /// Short names of all backend messages sent so far. + fn sent(&self) -> Vec<&'static str> { + self.sent + .lock() + .unwrap() + .iter() + .map(|m| match m { + PgWireBackendMessage::ParseComplete(_) => "ParseComplete", + PgWireBackendMessage::BindComplete(_) => "BindComplete", + PgWireBackendMessage::CloseComplete(_) => "CloseComplete", + PgWireBackendMessage::EmptyQueryResponse(_) => "EmptyQueryResponse", + PgWireBackendMessage::ParameterDescription(_) => "ParameterDescription", + PgWireBackendMessage::NoData(_) => "NoData", + PgWireBackendMessage::CommandComplete(_) => "CommandComplete", + PgWireBackendMessage::ReadyForQuery(_) => "ReadyForQuery", + _ => "other", + }) + .collect() + } + + /// Number of parameter types in the `ParameterDescription` at + /// `idx` among the sent messages. + fn parameter_description_len(&self, idx: usize) -> usize { + self.sent + .lock() + .unwrap() + .iter() + .filter_map(|m| match m { + PgWireBackendMessage::ParameterDescription(p) => Some(p.types.len()), + _ => None, + }) + .nth(idx) + .unwrap() + } + } + + impl ClientInfo for TestClient { + fn socket_addr(&self) -> SocketAddr { + self.inner.socket_addr() + } + + fn is_secure(&self) -> bool { + self.inner.is_secure() + } + + fn protocol_version(&self) -> crate::messages::ProtocolVersion { + self.inner.protocol_version() + } + + fn set_protocol_version(&mut self, version: crate::messages::ProtocolVersion) { + self.inner.set_protocol_version(version) + } + + fn pid_and_secret_key(&self) -> (i32, crate::messages::startup::SecretKey) { + self.inner.pid_and_secret_key() + } + + fn set_pid_and_secret_key( + &mut self, + pid: i32, + secret_key: crate::messages::startup::SecretKey, + ) { + self.inner.set_pid_and_secret_key(pid, secret_key) + } + + fn state(&self) -> PgWireConnectionState { + self.inner.state() + } + + fn set_state(&mut self, new_state: PgWireConnectionState) { + self.inner.set_state(new_state) + } + + fn transaction_status(&self) -> TransactionStatus { + self.inner.transaction_status() + } + + fn set_transaction_status(&mut self, new_status: TransactionStatus) { + self.inner.set_transaction_status(new_status) + } + + fn metadata(&self) -> &std::collections::HashMap { + self.inner.metadata() + } + + fn metadata_mut(&mut self) -> &mut std::collections::HashMap { + self.inner.metadata_mut() + } + + fn session_extensions(&self) -> &crate::api::SessionExtensions { + self.inner.session_extensions() + } + + #[cfg(any(feature = "_ring", feature = "_aws-lc-rs"))] + fn sni_server_name(&self) -> Option<&str> { + self.inner.sni_server_name() + } + + #[cfg(any(feature = "_ring", feature = "_aws-lc-rs"))] + fn client_certificates<'a>(&self) -> Option<&[rustls_pki_types::CertificateDer<'a>]> { + self.inner.client_certificates() + } + } + + impl ClientPortalStore for TestClient { + type PortalStore = crate::api::store::MemPortalStore; + + fn portal_store(&self) -> &Self::PortalStore { + self.inner.portal_store() + } + } + + impl Sink for TestClient { + type Error = PgWireError; + + fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn start_send(self: Pin<&mut Self>, item: PgWireBackendMessage) -> Result<(), Self::Error> { + self.sent.lock().unwrap().push(item); + Ok(()) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + /// A parser that fails on syntactically empty queries (they must never + /// reach a parser) and reports comment-only queries as empty. + #[derive(Default)] + struct RecordingParser { + calls: Mutex>, + } + + #[async_trait] + impl QueryParser for RecordingParser { + type Statement = String; + + async fn parse_sql( + &self, + _client: &C, + sql: &str, + _types: &[Option], + ) -> PgWireResult> + where + C: ClientInfo + Unpin + Send + Sync, + { + assert!( + !is_empty_query(sql), + "parser must never be called for an empty query, got {sql:?}" + ); + self.calls.lock().unwrap().push(sql.to_owned()); + if sql.starts_with("--") { + Ok(None) + } else { + Ok(Some(sql.to_owned())) + } + } + + fn get_parameter_types(&self, _stmt: &Self::Statement) -> PgWireResult> { + Ok(vec![]) + } + + fn get_result_schema( + &self, + _stmt: &Self::Statement, + _column_format: Option<&crate::api::portal::Format>, + ) -> PgWireResult> { + Ok(vec![]) + } + } + + struct TestHandler { + parser: Arc, + } + + impl TestHandler { + fn new() -> Self { + TestHandler { + parser: Arc::new(RecordingParser::default()), + } + } + } + + #[async_trait] + impl ExtendedQueryHandler for TestHandler { + type Statement = String; + type QueryParser = RecordingParser; + + fn query_parser(&self) -> Arc { + self.parser.clone() + } + + async fn do_query( + &self, + _client: &mut C, + _portal: &Portal, + _max_rows: usize, + ) -> PgWireResult + where + C: ClientInfo + ClientPortalStore + Sink + Unpin + Send + Sync, + C::PortalStore: PortalStore, + C::Error: Debug, + PgWireError: From<>::Error>, + { + Ok(Response::Execution(Tag::new("OK"))) + } + } + + fn parse(name: Option<&str>, query: &str) -> Parse { + Parse { + name: name.map(str::to_owned), + query: query.to_owned(), + type_oids: vec![], + } + } + + fn bind(portal: Option<&str>, statement: Option<&str>) -> Bind { + Bind { + portal_name: portal.map(str::to_owned), + statement_name: statement.map(str::to_owned), + parameter_format_codes: vec![], + parameters: vec![], + result_column_format_codes: vec![], + } + } + + fn describe(target_type: u8, name: Option<&str>) -> Describe { + Describe { + target_type, + name: name.map(str::to_owned), + } + } + + fn close(target_type: u8, name: Option<&str>) -> Close { + Close { + target_type, + name: name.map(str::to_owned), + } + } + + /// Parse/Bind/Describe/Execute/Sync of an empty query behaves exactly + /// like PostgreSQL. + #[tokio::test] + async fn empty_query_extended_protocol_sequence() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(None, "")) + .await + .unwrap(); + assert_eq!(client.sent(), ["ParseComplete"]); + + handler + ._on_describe(&mut client, describe(TARGET_TYPE_BYTE_STATEMENT, None)) + .await + .unwrap(); + assert_eq!(client.sent()[1..], ["ParameterDescription", "NoData"]); + assert_eq!(client.parameter_description_len(0), 0); + + handler + .on_bind(&mut client, bind(None, None)) + .await + .unwrap(); + assert_eq!(client.sent()[3..], ["BindComplete"]); + + handler + ._on_describe(&mut client, describe(TARGET_TYPE_BYTE_PORTAL, None)) + .await + .unwrap(); + assert_eq!(client.sent()[4..], ["NoData"]); + + // empty portals stay valid across repeated Execute + for _ in 0..2 { + handler + ._on_execute( + &mut client, + Execute { + name: None, + max_rows: 0, + }, + ) + .await + .unwrap(); + } + assert_eq!( + client.sent()[5..], + ["EmptyQueryResponse", "EmptyQueryResponse"] + ); + assert!(matches!( + client.state(), + PgWireConnectionState::ReadyForQuery + )); + + handler.on_sync(&mut client, PgSync).await.unwrap(); + assert_eq!(client.sent()[7..], ["ReadyForQuery"]); + + // Sync removed the unnamed empty portal + assert!(matches!( + handler + ._on_execute( + &mut client, + Execute { + name: None, + max_rows: 0 + }, + ) + .await, + Err(PgWireError::PortalNotFound(_)) + )); + + assert!(handler.parser.calls.lock().unwrap().is_empty()); + } + + /// A query the parser reports as empty (`None`) is stored and executed + /// as an empty query. + #[tokio::test] + async fn parser_reported_empty_query_executes_as_empty() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(Some("c"), "-- comment only")) + .await + .unwrap(); + assert_eq!( + handler.parser.calls.lock().unwrap().as_slice(), + ["-- comment only"] + ); + + handler + ._on_describe(&mut client, describe(TARGET_TYPE_BYTE_STATEMENT, Some("c"))) + .await + .unwrap(); + assert_eq!(client.sent()[1..], ["ParameterDescription", "NoData"]); + + handler + .on_bind(&mut client, bind(Some("p"), Some("c"))) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: Some("p".to_owned()), + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[3..], ["BindComplete", "EmptyQueryResponse"]); + } + + /// Semicolon-only and whitespace-only queries are empty in the extended + /// protocol as well. + #[tokio::test] + async fn semicolon_only_queries_are_empty() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(Some("s1"), ";")) + .await + .unwrap(); + handler + .on_parse(&mut client, parse(Some("s2"), " \n;\t")) + .await + .unwrap(); + assert_eq!(client.sent(), ["ParseComplete", "ParseComplete"]); + + handler + .on_bind(&mut client, bind(Some("p1"), Some("s1"))) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: Some("p1".to_owned()), + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[2..], ["BindComplete", "EmptyQueryResponse"]); + + assert!(handler.parser.calls.lock().unwrap().is_empty()); + } + + /// A string literal containing only a semicolon is a real query and must + /// reach the parser. + #[tokio::test] + async fn string_literal_semicolon_reaches_parser() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(None, "';'")) + .await + .unwrap(); + assert_eq!(client.sent(), ["ParseComplete"]); + assert_eq!(handler.parser.calls.lock().unwrap().as_slice(), ["';'"]); + + handler + .on_bind(&mut client, bind(None, None)) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: None, + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[1..], ["BindComplete", "CommandComplete"]); + } + + /// An empty Parse replaces a previously stored statement of the same + /// name, like PostgreSQL. + #[tokio::test] + async fn empty_parse_replaces_stored_statement() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(None, "select 1")) + .await + .unwrap(); + handler + .on_parse(&mut client, parse(None, "")) + .await + .unwrap(); + assert_eq!(client.sent(), ["ParseComplete", "ParseComplete"]); + + // describes the empty statement, not the previously parsed `select 1` + handler + ._on_describe(&mut client, describe(TARGET_TYPE_BYTE_STATEMENT, None)) + .await + .unwrap(); + assert_eq!(client.sent()[2..], ["ParameterDescription", "NoData"]); + + handler + .on_bind(&mut client, bind(None, None)) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: None, + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[4..], ["BindComplete", "EmptyQueryResponse"]); + } + + /// Binding an empty statement to a portal name replaces a previously + /// bound real portal of the same name. + #[tokio::test] + async fn empty_bind_replaces_stored_portal() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(None, "select 1")) + .await + .unwrap(); + handler + .on_bind(&mut client, bind(Some("p"), None)) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: Some("p".to_owned()), + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[1..], ["BindComplete", "CommandComplete"]); + + // re-parse the unnamed statement as empty, re-bind the same portal + handler + .on_parse(&mut client, parse(None, "")) + .await + .unwrap(); + handler + .on_bind(&mut client, bind(Some("p"), None)) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: Some("p".to_owned()), + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[4..], ["BindComplete", "EmptyQueryResponse"]); + } + + /// Bind on a statement that was never parsed is an error, and binding + /// parameters to an empty statement is a protocol violation (PostgreSQL + /// answers 08P01). + #[tokio::test] + async fn bind_errors() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + assert!(matches!( + handler + .on_bind(&mut client, bind(Some("p"), Some("missing"))) + .await, + Err(PgWireError::StatementNotFound(_)) + )); + + handler + .on_parse(&mut client, parse(Some("e"), "")) + .await + .unwrap(); + let mut with_param = bind(Some("p"), Some("e")); + with_param.parameters.push(Some(Bytes::from_static(b"1"))); + match handler.on_bind(&mut client, with_param).await { + Err(PgWireError::UserError(info)) => { + assert_eq!(info.code, "08P01"); + assert!(info.message.contains("requires 0")); + } + other => panic!("expected 08P01 user error, got {other:?}"), + } + } + + /// `Close` removes empty statements and empty portals like real ones. + #[tokio::test] + async fn close_removes_empty_statement_and_portal() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(Some("e"), "")) + .await + .unwrap(); + handler + .on_close(&mut client, close(TARGET_TYPE_BYTE_STATEMENT, Some("e"))) + .await + .unwrap(); + assert_eq!(client.sent(), ["ParseComplete", "CloseComplete"]); + assert!(matches!( + handler + .on_bind(&mut client, bind(Some("p"), Some("e"))) + .await, + Err(PgWireError::StatementNotFound(_)) + )); + + handler + .on_parse(&mut client, parse(Some("e2"), "")) + .await + .unwrap(); + handler + .on_bind(&mut client, bind(Some("p2"), Some("e2"))) + .await + .unwrap(); + handler + .on_close(&mut client, close(TARGET_TYPE_BYTE_PORTAL, Some("p2"))) + .await + .unwrap(); + assert!(matches!( + handler + ._on_execute( + &mut client, + Execute { + name: Some("p2".to_owned()), + max_rows: 0, + }, + ) + .await, + Err(PgWireError::PortalNotFound(_)) + )); + } + + /// After a failed Bind the empty statement stays prepared, matching + /// PostgreSQL. + #[tokio::test] + async fn marker_survives_failed_bind() { + let handler = TestHandler::new(); + let mut client = TestClient::new(); + + handler + .on_parse(&mut client, parse(Some("e"), "")) + .await + .unwrap(); + let mut bad = bind(Some("p"), Some("e")); + bad.parameters.push(Some(Bytes::from_static(b"1"))); + assert!(handler.on_bind(&mut client, bad).await.is_err()); + + handler + .on_bind(&mut client, bind(Some("p"), Some("e"))) + .await + .unwrap(); + handler + ._on_execute( + &mut client, + Execute { + name: Some("p".to_owned()), + max_rows: 0, + }, + ) + .await + .unwrap(); + assert_eq!(client.sent()[1..], ["BindComplete", "EmptyQueryResponse"]); + } +} diff --git a/src/api/stmt.rs b/src/api/stmt.rs index 3cefa65..1bec915 100644 --- a/src/api/stmt.rs +++ b/src/api/stmt.rs @@ -9,6 +9,7 @@ use crate::messages::PgWireBackendMessage; use crate::messages::extendedquery::Parse; use super::portal::Format; +use super::query::is_empty_query; use super::results::FieldInfo; use super::{ClientInfo, DEFAULT_NAME}; @@ -26,12 +27,15 @@ pub struct StoredStatement { } impl StoredStatement { - /// Parse a `Parse` message into a stored statement using the given query parser. + /// Parse a `Parse` message into a stored statement using the given query + /// parser. + /// + /// Returns `None` for an empty query: there is no statement to store. pub async fn parse( client: &C, parse: &Parse, parser: Q, - ) -> PgWireResult> + ) -> PgWireResult>> where C: ClientInfo + Sink + Unpin + Send + Sync, Q: QueryParser, @@ -41,15 +45,20 @@ impl StoredStatement { .iter() .map(|oid| Type::from_oid(*oid)) .collect::>(); - let statement = parser.parse_sql(client, &parse.query, &types).await?; - Ok(StoredStatement { - id: parse - .name - .clone() - .unwrap_or_else(|| DEFAULT_NAME.to_owned()), - statement, - parameter_types: types, - }) + if is_empty_query(&parse.query) { + return Ok(None); + } + Ok(parser + .parse_sql(client, &parse.query, &types) + .await? + .map(|statement| StoredStatement { + id: parse + .name + .clone() + .unwrap_or_else(|| DEFAULT_NAME.to_owned()), + statement, + parameter_types: types, + })) } } @@ -63,12 +72,17 @@ pub trait QueryParser { /// /// The client may or may not provide type information with any parameters /// from the sql. + /// + /// Return `Ok(None)` for an empty query; it is stored as an empty + /// statement and executes to `EmptyQueryResponse`, like PostgreSQL. + /// Syntactically empty queries (only semicolons and whitespace) never + /// reach this method. async fn parse_sql( &self, client: &C, sql: &str, types: &[Option], - ) -> PgWireResult + ) -> PgWireResult> where C: ClientInfo + Unpin + Send + Sync; @@ -103,7 +117,7 @@ where client: &C, sql: &str, types: &[Option], - ) -> PgWireResult + ) -> PgWireResult> where C: ClientInfo + Unpin + Send + Sync, { @@ -136,11 +150,11 @@ impl QueryParser for NoopQueryParser { _client: &C, sql: &str, _types: &[Option], - ) -> PgWireResult + ) -> PgWireResult> where C: ClientInfo + Unpin + Send + Sync, { - Ok(sql.to_owned()) + Ok(Some(sql.to_owned())) } fn get_parameter_types(&self, _stmt: &Self::Statement) -> PgWireResult> { diff --git a/src/api/store.rs b/src/api/store.rs index f158262..e51439c 100644 --- a/src/api/store.rs +++ b/src/api/store.rs @@ -5,7 +5,45 @@ use std::sync::{Arc, RwLock}; use super::portal::Portal; use super::stmt::StoredStatement; +/// An entry stored in a [`PortalStore`] under a statement or portal name: +/// either the stored value, or an empty marker for a query that parsed to +/// no statement. +#[derive(Debug)] +pub enum Entry { + Empty, + Value(Arc), +} + +impl Clone for Entry { + fn clone(&self) -> Self { + match self { + Entry::Empty => Entry::Empty, + Entry::Value(value) => Entry::Value(Arc::clone(value)), + } + } +} + +impl Entry { + /// The stored value, if any. + pub fn value(&self) -> Option<&Arc> { + match self { + Entry::Empty => None, + Entry::Value(value) => Some(value), + } + } + + /// Whether this is an empty entry. + pub fn is_empty(&self) -> bool { + matches!(self, Entry::Empty) + } +} + /// Storage trait for prepared statements and portals. +/// +/// Statements `Parse`d from empty queries and portals bound from them are +/// stored as empty entries, like PostgreSQL: every `put_*` replaces whatever +/// was previously stored under the name, and `rm_*`/`clear_portals` remove +/// empty entries along with regular ones. pub trait PortalStore: Any + Send + Sync + 'static { type Statement; @@ -15,15 +53,21 @@ pub trait PortalStore: Any + Send + Sync + 'static { /// Store a prepared statement by name. fn put_statement(&self, statement: Arc>); + /// Store an empty prepared statement by name. + fn put_empty_statement(&self, name: &str); + /// Remove a prepared statement by name. fn rm_statement(&self, name: &str); /// Retrieve a prepared statement by name. - fn get_statement(&self, name: &str) -> Option>>; + fn get_statement(&self, name: &str) -> Option>>; /// Store a portal by name. fn put_portal(&self, portal: Arc>); + /// Store an empty portal by name. + fn put_empty_portal(&self, name: &str); + /// Remove a portal by name. fn rm_portal(&self, name: &str); @@ -31,16 +75,16 @@ pub trait PortalStore: Any + Send + Sync + 'static { fn clear_portals(&self); /// Retrieve a portal by name. - fn get_portal(&self, name: &str) -> Option>>; + fn get_portal(&self, name: &str) -> Option>>; } /// In-memory implementation of `PortalStore` backed by `BTreeMap`. #[derive(Debug, Default, new)] pub struct MemPortalStore { #[new(default)] - statements: RwLock>>>, + statements: RwLock>>>, #[new(default)] - portals: RwLock>>>, + portals: RwLock>>>, } impl PortalStore for MemPortalStore { @@ -51,8 +95,14 @@ impl PortalStore for MemPortalStore { } fn put_statement(&self, statement: Arc>) { + let name = statement.id.to_owned(); let mut guard = self.statements.write().unwrap(); - guard.insert(statement.id.to_owned(), statement); + guard.insert(name, Entry::Value(statement)); + } + + fn put_empty_statement(&self, name: &str) { + let mut guard = self.statements.write().unwrap(); + guard.insert(name.to_owned(), Entry::Empty); } fn rm_statement(&self, name: &str) { @@ -60,14 +110,19 @@ impl PortalStore for MemPortalStore { guard.remove(name); } - fn get_statement(&self, name: &str) -> Option>> { + fn get_statement(&self, name: &str) -> Option>> { let guard = self.statements.read().unwrap(); guard.get(name).cloned() } fn put_portal(&self, portal: Arc>) { let mut guard = self.portals.write().unwrap(); - guard.insert(portal.name.to_owned(), portal); + guard.insert(portal.name.to_owned(), Entry::Value(portal)); + } + + fn put_empty_portal(&self, name: &str) { + let mut guard = self.portals.write().unwrap(); + guard.insert(name.to_owned(), Entry::Empty); } fn rm_portal(&self, name: &str) { @@ -80,8 +135,62 @@ impl PortalStore for MemPortalStore { guard.clear(); } - fn get_portal(&self, name: &str) -> Option>> { + fn get_portal(&self, name: &str) -> Option>> { let guard = self.portals.read().unwrap(); guard.get(name).cloned() } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn statement_entries_replace_each_other() { + let store: MemPortalStore = MemPortalStore::new(); + assert!(store.get_statement("s").is_none()); + + store.put_empty_statement("s"); + assert!(store.get_statement("s").unwrap().is_empty()); + + store.put_statement(Arc::new(StoredStatement::new( + "s".to_owned(), + "select 1".to_owned(), + vec![], + ))); + assert_eq!( + store + .get_statement("s") + .and_then(|e| e.value().map(|s| s.statement.clone())), + Some("select 1".to_owned()) + ); + + store.put_empty_statement("s"); + assert!(store.get_statement("s").unwrap().is_empty()); + + store.rm_statement("s"); + assert!(store.get_statement("s").is_none()); + } + + #[test] + fn portal_entries_replace_each_other_and_clear() { + let store: MemPortalStore = MemPortalStore::new(); + let statement = Arc::new(StoredStatement::new( + "s".to_owned(), + "select 1".to_owned(), + vec![], + )); + let portal = Portal::new_cursor("p".to_owned(), statement); + + store.put_portal(Arc::new(portal)); + assert!(store.get_portal("p").unwrap().value().is_some()); + + store.put_empty_portal("p"); + assert!(store.get_portal("p").unwrap().is_empty()); + + store.put_empty_portal("p2"); + store.clear_portals(); + assert!(store.get_portal("p").is_none()); + assert!(store.get_portal("p2").is_none()); + } +}