use std::fmt;
use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PgStoreConfig {
max_connections: u32,
min_connections: u32,
acquire_timeout: Duration,
idle_timeout: Option<Duration>,
max_lifetime: Option<Duration>,
statement_timeout: Option<Duration>,
schema: Option<String>,
}
impl PgStoreConfig {
pub const DEFAULT_MAX_CONNECTIONS: u32 = 10;
pub const DEFAULT_MIN_CONNECTIONS: u32 = 1;
pub const DEFAULT_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(5);
pub const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(600);
pub const DEFAULT_MAX_LIFETIME: Duration = Duration::from_secs(1800);
#[must_use]
pub fn new() -> Self {
Self {
max_connections: Self::DEFAULT_MAX_CONNECTIONS,
min_connections: Self::DEFAULT_MIN_CONNECTIONS,
acquire_timeout: Self::DEFAULT_ACQUIRE_TIMEOUT,
idle_timeout: Some(Self::DEFAULT_IDLE_TIMEOUT),
max_lifetime: Some(Self::DEFAULT_MAX_LIFETIME),
statement_timeout: None,
schema: None,
}
}
#[must_use]
pub fn max_connections(mut self, connections: u32) -> Self {
self.max_connections = connections;
self
}
#[must_use]
pub fn min_connections(mut self, connections: u32) -> Self {
self.min_connections = connections;
self
}
#[must_use]
pub fn acquire_timeout(mut self, timeout: Duration) -> Self {
self.acquire_timeout = timeout;
self
}
#[must_use]
pub fn idle_timeout(mut self, timeout: Option<Duration>) -> Self {
self.idle_timeout = timeout;
self
}
#[must_use]
pub fn max_lifetime(mut self, lifetime: Option<Duration>) -> Self {
self.max_lifetime = lifetime;
self
}
#[must_use]
pub fn statement_timeout(mut self, timeout: Duration) -> Self {
self.statement_timeout = Some(timeout);
self
}
#[must_use]
pub fn inherit_statement_timeout(mut self) -> Self {
self.statement_timeout = None;
self
}
pub fn schema(mut self, schema: impl Into<String>) -> Result<Self, ConfigError> {
let schema = schema.into();
if !is_plain_identifier(&schema) {
return Err(ConfigError::InvalidSchemaName);
}
self.schema = Some(schema);
Ok(self)
}
#[must_use]
pub fn schema_name(&self) -> Option<&str> {
self.schema.as_deref()
}
#[must_use]
pub fn statement_timeout_value(&self) -> Option<Duration> {
self.statement_timeout
}
#[must_use]
pub fn connection_bounds(&self) -> (u32, u32) {
(self.min_connections, self.max_connections)
}
pub(crate) fn apply_pool(
&self,
options: sqlx::postgres::PgPoolOptions,
) -> sqlx::postgres::PgPoolOptions {
options
.max_connections(self.max_connections)
.min_connections(self.min_connections)
.acquire_timeout(self.acquire_timeout)
.idle_timeout(self.idle_timeout)
.max_lifetime(self.max_lifetime)
}
pub(crate) fn apply_connection(
&self,
mut options: sqlx::postgres::PgConnectOptions,
) -> sqlx::postgres::PgConnectOptions {
if let Some(schema) = &self.schema {
options = options.options([("search_path", schema.as_str())]);
}
if let Some(timeout) = self.statement_timeout {
let millis = timeout.as_millis().to_string();
options = options.options([("statement_timeout", millis.as_str())]);
}
options
}
}
impl Default for PgStoreConfig {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum ConfigError {
#[error("schema name is not a plain unquoted PostgreSQL identifier")]
InvalidSchemaName,
}
fn is_plain_identifier(name: &str) -> bool {
const MAX_IDENTIFIER_BYTES: usize = 63;
let mut characters = name.chars();
let Some(first) = characters.next() else {
return false;
};
name.len() <= MAX_IDENTIFIER_BYTES
&& (first.is_ascii_lowercase() || first.is_ascii_uppercase() || first == '_')
&& characters.all(|character| character.is_ascii_alphanumeric() || character == '_')
}
impl fmt::Display for PgStoreConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"pool {}..{}, schema {}",
self.min_connections,
self.max_connections,
self.schema.as_deref().unwrap_or("<default>")
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn defaults_are_the_documented_ones() {
let config = PgStoreConfig::new();
assert_eq!(config.connection_bounds(), (1, 10));
assert_eq!(config.schema_name(), None);
assert_eq!(config.statement_timeout_value(), None);
assert_eq!(config, PgStoreConfig::default());
}
#[test]
fn a_schema_name_must_be_a_plain_identifier() {
assert!(PgStoreConfig::new().schema("turnframe").is_ok());
assert!(PgStoreConfig::new().schema("_tf_test_1").is_ok());
for rejected in [
"",
"1leading_digit",
"has space",
"quote\"injection",
"semicolon;drop",
"dotted.name",
"unicodé",
] {
assert_eq!(
PgStoreConfig::new().schema(rejected).unwrap_err(),
ConfigError::InvalidSchemaName,
"{rejected:?} must be refused"
);
}
let too_long = "a".repeat(64);
assert_eq!(
PgStoreConfig::new().schema(too_long).unwrap_err(),
ConfigError::InvalidSchemaName
);
}
#[test]
fn display_never_carries_a_connection_string() {
let config = PgStoreConfig::new().schema("turnframe").unwrap();
assert_eq!(config.to_string(), "pool 1..10, schema turnframe");
assert_eq!(
PgStoreConfig::new().to_string(),
"pool 1..10, schema <default>"
);
}
#[test]
fn timeouts_are_settable_and_clearable() {
let config = PgStoreConfig::new()
.statement_timeout(Duration::from_secs(3))
.acquire_timeout(Duration::from_secs(1))
.idle_timeout(None)
.max_lifetime(None)
.min_connections(2)
.max_connections(4);
assert_eq!(
config.statement_timeout_value(),
Some(Duration::from_secs(3))
);
assert_eq!(config.connection_bounds(), (2, 4));
assert_eq!(
config.inherit_statement_timeout().statement_timeout_value(),
None
);
}
}