diff --git a/integration/unix_sockets/pgdog.toml b/integration/unix_sockets/pgdog.toml new file mode 100644 index 000000000..601b937f1 --- /dev/null +++ b/integration/unix_sockets/pgdog.toml @@ -0,0 +1,15 @@ +# ------------------------------------------------------------------------------ +# ----- General ---------------------------------------------------------------- + +[general] +auth_type = "trust" + +# ------------------------------------------------------------------------------ +# ----- Database :: pgdog ------------------------------------------------------ + +[[databases]] +name = "pgdog" +host = "/tmp" +port = 5432 +database_name = "pgdog" +user = "pgdog" diff --git a/integration/unix_sockets/run.sh b/integration/unix_sockets/run.sh new file mode 100755 index 000000000..342550451 --- /dev/null +++ b/integration/unix_sockets/run.sh @@ -0,0 +1,52 @@ +#!/bin/bash +# End-to-end test: pgdog connecting to Postgres over a Unix domain socket. +set -euo pipefail +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd ) +source ${SCRIPT_DIR}/../common.sh + +export PGPASSWORD=pgdog +CTL_PSQL=(psql -h 127.0.0.1 -p 5432 -U pgdog -d postgres -t -A) + +# Detect socket dir + port from the running Postgres. +SOCKET_DIR=$("${CTL_PSQL[@]}" -c "show unix_socket_directories" | cut -d, -f1) +PG_BACKEND_PORT=$("${CTL_PSQL[@]}" -c "show port") +if [ -z "${SOCKET_DIR}" ]; then + echo "FAIL: could not detect unix_socket_directories from Postgres" >&2 + exit 1 +fi +echo "Postgres unix socket directory: ${SOCKET_DIR} (port ${PG_BACKEND_PORT})" + +# Pre-flight: pgdog must be able to reach Postgres over the socket as pgdog. +if ! psql -h "${SOCKET_DIR}" -p "${PG_BACKEND_PORT}" -U pgdog -d postgres \ + -c "select 1" >/dev/null 2>&1; then + echo "FAIL: cannot connect to Postgres over unix socket ${SOCKET_DIR} as pgdog." >&2 + echo " Check pg_hba.conf 'local' lines: trust, or peer with OS user pgdog." >&2 + exit 1 +fi + +# Patch the static config with the detected socket dir/port into a temp dir. +TMP_CFG_DIR=$(mktemp -d /tmp/pgdog-unix-cfg.XXXXXX) +sed -e "s|^host = .*|host = \"${SOCKET_DIR}\"|" \ + -e "s|^port = .*|port = ${PG_BACKEND_PORT}|" \ + "${SCRIPT_DIR}/pgdog.toml" > "${TMP_CFG_DIR}/pgdog.toml" +cp "${SCRIPT_DIR}/users.toml" "${TMP_CFG_DIR}/" + +run_pgdog "${TMP_CFG_DIR}" +wait_for_pgdog + +# 1. Query through the proxy. +psql -h 127.0.0.1 -p 6432 -U pgdog -d pgdog -v ON_ERROR_STOP=1 \ + -c "select version()" >/dev/null +echo "PASS: query through pgdog" + +# 2. Backend connections over the Unix socket have client_addr IS NULL. +CONNS=$("${CTL_PSQL[@]}" -c \ + "select count(*) from pg_stat_activity where usename = 'pgdog' and backend_type = 'client backend' and client_addr is null") +if [ -z "${CONNS}" ] || [ "${CONNS}" = "0" ]; then + echo "FAIL: no backend connections over Unix socket (client_addr IS NULL)" >&2 + exit 1 +fi +echo "PASS: ${CONNS} backend connection(s) over Unix socket" + +stop_pgdog +rm -rf "${TMP_CFG_DIR}" diff --git a/integration/unix_sockets/users.toml b/integration/unix_sockets/users.toml new file mode 100644 index 000000000..581cdb75b --- /dev/null +++ b/integration/unix_sockets/users.toml @@ -0,0 +1,4 @@ +[[users]] +name = "pgdog" +database = "pgdog" +password = "pgdog" diff --git a/pgdog/src/admin/show_bans.rs b/pgdog/src/admin/show_bans.rs index 91eb75272..963a8f72a 100644 --- a/pgdog/src/admin/show_bans.rs +++ b/pgdog/src/admin/show_bans.rs @@ -55,7 +55,7 @@ impl Command for ShowBans { row.add(pool.id() as i64) .add(user.database.as_str()) .add(user.user.as_str()) - .add(pool.addr().host.as_str()) + .add(pool.addr().host.to_string()) .add(pool.addr().port as i64) .add(shard_num as i64) .add(role.to_string()) diff --git a/pgdog/src/admin/show_pools.rs b/pgdog/src/admin/show_pools.rs index 5366dc5c4..9a68993c5 100644 --- a/pgdog/src/admin/show_pools.rs +++ b/pgdog/src/admin/show_pools.rs @@ -58,7 +58,7 @@ impl Command for ShowPools { row.add(pool.id() as i64) .add(user.database.as_str()) .add(user.user.as_str()) - .add(pool.addr().host.as_str()) + .add(pool.addr().host.to_string()) .add(pool.addr().port as i64) .add(shard_num as i64) .add(role.to_string()) diff --git a/pgdog/src/admin/show_replication.rs b/pgdog/src/admin/show_replication.rs index 3a8ca15f2..68190e6cc 100644 --- a/pgdog/src/admin/show_replication.rs +++ b/pgdog/src/admin/show_replication.rs @@ -53,7 +53,7 @@ impl Command for ShowReplication { row.add(pool.id() as i64) .add(user.database.as_str()) .add(user.user.as_str()) - .add(pool.addr().host.as_str()) + .add(pool.addr().host.to_string()) .add(pool.addr().port as i64) .add(shard_num as i64) .add(role.to_string()) diff --git a/pgdog/src/admin/show_server_memory.rs b/pgdog/src/admin/show_server_memory.rs index 88f81ae50..bcb3915fb 100644 --- a/pgdog/src/admin/show_server_memory.rs +++ b/pgdog/src/admin/show_server_memory.rs @@ -43,7 +43,7 @@ impl Command for ShowServerMemory { row.add(server.stats.pool_id as i64) .add(server.addr.database_name.as_str()) .add(server.addr.user.as_str()) - .add(server.addr.host.as_str()) + .add(server.addr.host.to_string()) .add(server.addr.port as i64) .add(server.stats.id) .add(memory.buffer.reallocs as i64) diff --git a/pgdog/src/admin/show_servers.rs b/pgdog/src/admin/show_servers.rs index 287db00c6..d6dab7fe2 100644 --- a/pgdog/src/admin/show_servers.rs +++ b/pgdog/src/admin/show_servers.rs @@ -90,7 +90,7 @@ impl Command for ShowServers { .add("pool_id", server.stats.pool_id) .add("database", server.addr.database_name) .add("user", server.addr.user) - .add("addr", server.addr.host.as_str()) + .add("addr", server.addr.host.to_string()) .add("port", server.addr.port.to_string()) .add("state", server.stats.state.to_string()) .add( diff --git a/pgdog/src/admin/show_stats.rs b/pgdog/src/admin/show_stats.rs index 740095edc..7f253ede3 100644 --- a/pgdog/src/admin/show_stats.rs +++ b/pgdog/src/admin/show_stats.rs @@ -78,7 +78,7 @@ impl Command for ShowStats { dr.add(user.database.as_str()) .add(user.user.as_str()) - .add(&pool.addr().host) + .add(pool.addr().host.to_string()) .add(pool.addr().port as i64) .add(shard_num) .add(role.to_string()); diff --git a/pgdog/src/backend/auth/azure_workload_identity.rs b/pgdog/src/backend/auth/azure_workload_identity.rs index 6346b2123..9f25dbac9 100644 --- a/pgdog/src/backend/auth/azure_workload_identity.rs +++ b/pgdog/src/backend/auth/azure_workload_identity.rs @@ -40,6 +40,7 @@ mod tests { use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; use super::*; + use crate::backend::pool::transport::Transport; use crate::config::ServerAuth; use crate::test_utils::set_env_var; use pgdog_config::Role; diff --git a/pgdog/src/backend/auth/rds_iam.rs b/pgdog/src/backend/auth/rds_iam.rs index 8e3d3aa76..cbef7ef2e 100644 --- a/pgdog/src/backend/auth/rds_iam.rs +++ b/pgdog/src/backend/auth/rds_iam.rs @@ -37,7 +37,8 @@ fn resolve_region(addr: &Address) -> Result { return Ok(region.clone()); } - infer_region_from_rds_host(&addr.host).ok_or_else(|| { + let host = addr.host.tcp()?; + infer_region_from_rds_host(host).ok_or_else(|| { Error::RdsIamToken(format!( "unable to infer AWS region from host \"{}\"; set \"server_iam_region\"", addr.host @@ -53,9 +54,10 @@ fn resolve_region(addr: &Address) -> Result { pub(crate) async fn token(addr: Address) -> Result<(String, SystemTime), Error> { let region = resolve_region(&addr)?; let sdk_config = aws_config::load_defaults(BehaviorVersion::latest()).await; + let host = addr.host.tcp()?; let config = AuthTokenConfig::builder() - .hostname(addr.host.as_str()) + .hostname(host) .port(addr.port.into()) .username(addr.user.as_str()) .region(Region::new(region.clone())) @@ -88,6 +90,7 @@ mod tests { use pgdog_config::Role; use super::*; + use crate::backend::pool::transport::Transport; use crate::config::ServerAuth; use crate::test_utils::set_env_var; diff --git a/pgdog/src/backend/auth/vault.rs b/pgdog/src/backend/auth/vault.rs index 297a74af0..102070281 100644 --- a/pgdog/src/backend/auth/vault.rs +++ b/pgdog/src/backend/auth/vault.rs @@ -151,6 +151,7 @@ mod tests { use super::*; use crate::auth::vault::{VAULT_TOKEN, VaultToken}; + use crate::backend::pool::transport::Transport; use crate::config::ConfigAndUsers; fn setup() { diff --git a/pgdog/src/backend/error.rs b/pgdog/src/backend/error.rs index 6b344c820..01f10bb5d 100644 --- a/pgdog/src/backend/error.rs +++ b/pgdog/src/backend/error.rs @@ -1,6 +1,6 @@ use thiserror::Error; -use crate::net::messages::ErrorResponse; +use crate::{backend::pool::transport::TransportError, net::messages::ErrorResponse}; use super::databases::User; @@ -146,6 +146,9 @@ pub enum Error { #[error("missing canonical oid for type {0}")] MissingCanonicalOid(String), + + #[error(transparent)] + Transport(#[from] TransportError), } impl From for Error { diff --git a/pgdog/src/backend/pool/address.rs b/pgdog/src/backend/pool/address.rs index 9443a9069..cb65bdf94 100644 --- a/pgdog/src/backend/pool/address.rs +++ b/pgdog/src/backend/pool/address.rs @@ -12,13 +12,14 @@ use crate::backend::Error; use crate::backend::auth::{azure_workload_identity, rds_iam, vault}; use crate::backend::pool::dns_cache::DnsCache; use crate::backend::pool::token_cache::TokenCache; +use crate::backend::pool::transport::{Transport, unix_socket_path}; use crate::config::{Database, ServerAuth, User, config}; /// Server address. -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default, Eq, Hash)] +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)] pub struct Address { /// Server host. - pub host: String, + pub host: Transport, /// Server port. pub port: u16, /// PostgreSQL database name. @@ -45,10 +46,29 @@ pub struct Address { pub configured_role: Role, } +impl Default for Address { + /// Local development defaults: `pgdog` on `127.0.0.1:5432`. ` + fn default() -> Self { + Address { + host: Transport::TCP("127.0.0.1".to_string()), + port: 5432, + user: "pgdog".into(), + passwords: vec!["pgdog".into()], + database_name: "pgdog".into(), + server_auth: ServerAuth::Password, + server_iam_region: None, + vault_path: None, + vault_refresh_percent: None, + database_number: 0, + configured_role: Role::Primary, + } + } +} + impl From
for pgdog_stats::Address { fn from(value: Address) -> Self { pgdog_stats::Address { - host: value.host, + host: value.host.to_string(), port: value.port, database_name: value.database_name, user: value.user, @@ -64,9 +84,10 @@ impl Address { /// Create new address from config values. pub(crate) fn new(database: &Database, user: &User, database_number: usize) -> Self { let server_auth = user.server_auth; + let host = database.host.clone(); Address { - host: database.host.clone(), + host: host.into(), port: database.port, database_name: if let Some(database_name) = database.database_name.clone() { database_name @@ -168,9 +189,10 @@ impl Address { /// pub(crate) async fn addr(&self) -> Result { let dns_cache_override_enabled = config().config.general.dns_ttl().is_some(); + let host = self.host.tcp()?; if dns_cache_override_enabled { - let ip = DnsCache::global().resolve(&self.host).await?; + let ip = DnsCache::global().resolve(host).await?; return Ok(SocketAddr::new(ip, self.port)); } @@ -179,7 +201,7 @@ impl Address { socket_addrs .next() - .ok_or(Error::DnsResolutionFailed(self.host.clone())) + .ok_or(Error::DnsResolutionFailed(host.to_string())) } /// A replacement for [`PartialEq`] which accounts for @@ -196,34 +218,23 @@ impl Address { true } } - - /// Test convention: `new_test()` represents a primary. Tests that need - /// a replica do `Address { configured_role: Role::Replica, ..new_test() }`. - #[cfg(test)] - pub fn new_test() -> Self { - Self { - host: "127.0.0.1".into(), - port: 5432, - user: "pgdog".into(), - passwords: vec!["pgdog".into()], - database_name: "pgdog".into(), - server_auth: ServerAuth::Password, - server_iam_region: None, - vault_path: None, - vault_refresh_percent: None, - database_number: 0, - configured_role: Role::Primary, - } - } } impl std::fmt::Display for Address { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!( - f, - "{}@{}:{}/{}", - self.user, self.host, self.port, self.database_name - ) + match &self.host { + Transport::TCP(host) => { + write!( + f, + "{}@{}:{}/{}", + self.user, host, self.port, self.database_name + ) + } + Transport::Unix(dir) => { + let file = unix_socket_path(dir, &self.port); + write!(f, "{}", file.display()) + } + } } } @@ -237,17 +248,21 @@ impl TryFrom for Address { let password = value.password().ok_or(())?.to_string(); let database_name = value.path().replace("/", "").to_string(); - // A URL says nothing about role; fall through to `Role::Auto` - // via the derived `Default`. The PROBE command (the only caller) - // never reads `configured_role` anyway. + // A URL says nothing about role; fall through to the explicit + // `Role::Auto` below. The PROBE command (the only caller) never + // reads `configured_role` anyway. Ok(Self { - host, + host: host.into(), port, passwords: vec![password.into()], user, database_name, server_auth: ServerAuth::Password, - ..Default::default() + server_iam_region: None, + vault_path: None, + vault_refresh_percent: None, + database_number: 0, + configured_role: Role::Auto, }) } } @@ -256,7 +271,7 @@ impl TryFrom for Address { mod test { use std::time::{Duration, Instant, SystemTime}; - use crate::config; + use crate::{backend::pool::transport::unix_socket_path, config}; use super::*; @@ -280,7 +295,7 @@ mod test { let address = Address::new(&database, &user, 0); - assert_eq!(address.host, "127.0.0.1"); + assert_eq!(address.host.to_string(), "127.0.0.1"); assert_eq!(address.port, 6432); assert_eq!(address.database_name, "pgdog"); assert_eq!(address.user, "pgdog"); @@ -361,7 +376,7 @@ mod test { let addr = Address::try_from(Url::parse("postgres://user:password@127.0.0.1:6432/pgdb").unwrap()) .unwrap(); - assert_eq!(addr.host, "127.0.0.1"); + assert_eq!(addr.host.to_string(), "127.0.0.1"); assert_eq!(addr.port, 6432); assert_eq!(addr.database_name, "pgdb"); assert_eq!(addr.user, "user"); @@ -372,7 +387,7 @@ mod test { #[test] fn test_compatible_ignores_password_changes() { - let address = Address::new_test(); + let address = Address::default(); let mut rotated = address.clone(); rotated.passwords = vec!["rotated".into()]; @@ -381,7 +396,7 @@ mod test { #[test] fn test_compatible_rejects_other_field_changes() { - let address = Address::new_test(); + let address = Address::default(); let mut moved = address.clone(); moved.port += 1; @@ -392,13 +407,13 @@ mod test { #[tokio::test] async fn test_auth_secret_password_mode() { - let addr = Address::new_test(); + let addr = Address::default(); assert_eq!(addr.auth_secrets().await.unwrap().first().unwrap(), "pgdog"); } #[tokio::test] async fn test_auth_secrets_returns_valid_password_first() { - let mut addr = Address::new_test(); + let mut addr = Address::default(); let invalid1: Password = "invalid1".into(); let invalid2: Password = "invalid2".into(); let valid: Password = "valid".into(); @@ -412,7 +427,7 @@ mod test { assert!(secrets.first().unwrap().is_valid()); // Even if the valid password is last, it should still come first. - let mut addr = Address::new_test(); + let mut addr = Address::default(); let invalid1: Password = "invalid1".into(); let invalid2: Password = "invalid2".into(); let valid: Password = "valid".into(); @@ -424,7 +439,7 @@ mod test { assert_eq!(secrets.first().unwrap(), "valid"); // With multiple valid passwords, a valid one is still first. - let mut addr = Address::new_test(); + let mut addr = Address::default(); let invalid: Password = "invalid".into(); invalid.valid(false); addr.passwords = vec![invalid, "valid_a".into(), "valid_b".into()]; @@ -435,7 +450,7 @@ mod test { assert!(head == "valid_a" || head == "valid_b"); // Flipping validity at runtime changes which password comes first. - let mut addr = Address::new_test(); + let mut addr = Address::default(); let first: Password = "first".into(); let second: Password = "second".into(); addr.passwords = vec![first.clone(), second.clone()]; @@ -565,7 +580,7 @@ mod test { #[tokio::test] async fn test_auth_credentials_password_mode_uses_configured_user() { - let addr = Address::new_test(); + let addr = Address::default(); let (user, secrets) = addr.auth_credentials().await.unwrap(); assert_eq!(user, "pgdog"); assert_eq!(secrets.first().unwrap(), "pgdog"); @@ -697,7 +712,7 @@ mod test { let addr = Address { host: hostname.into(), port: 15432, - ..Address::new_test() + ..Default::default() }; let socket_addr = addr.addr().await.expect("resolve address"); @@ -708,4 +723,32 @@ mod test { Some(socket_addr.ip()) ); } + + #[test] + fn test_transport_parsing() { + // Absolute paths are Unix socket directories. + assert!(matches!(Transport::new("/tmp"), Transport::Unix(_))); + assert!(matches!( + Transport::new("/var/run/postgresql"), + Transport::Unix(_) + )); + // Everything else is a TCP hostname. + assert!(matches!(Transport::new("127.0.0.1"), Transport::TCP(_))); + assert!(matches!(Transport::new("localhost"), Transport::TCP(_))); + } + + #[test] + fn test_unix_socket_path() { + let dir = std::path::PathBuf::from("/tmp"); + assert_eq!( + unix_socket_path(&dir, &5432), + std::path::PathBuf::from("/tmp/.s.PGSQL.5432") + ); + + // Any port formats into the socket file name. + assert_eq!( + unix_socket_path(&dir, &54321), + std::path::PathBuf::from("/tmp/.s.PGSQL.54321") + ); + } } diff --git a/pgdog/src/backend/pool/cluster.rs b/pgdog/src/backend/pool/cluster.rs index d10dfb85d..fc03323cc 100644 --- a/pgdog/src/backend/pool/cluster.rs +++ b/pgdog/src/backend/pool/cluster.rs @@ -814,13 +814,13 @@ mod test { database: "pgdog".into(), }); let primary = Some(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config: Config::default(), }); let replicas = &[PoolConfig { address: Address { configured_role: Role::Replica, - ..Address::new_test() + ..Default::default() }, config: Config::default(), }]; @@ -959,7 +959,7 @@ mod test { primary: Some(&PoolConfig { address: Address { database_name: "pgdog1".into(), - ..Address::new_test() + ..Default::default() }, config: Config::default(), }), @@ -967,7 +967,7 @@ mod test { address: Address { database_name: "pgdog1".into(), configured_role: Role::Replica, - ..Address::new_test() + ..Default::default() }, config: Config::default(), }], @@ -987,7 +987,7 @@ mod test { Cluster { shards: vec![Shard::new(ShardConfig { primary: Some(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config: Config::default(), }), identifier: identifier.clone(), @@ -1018,7 +1018,7 @@ mod test { replicas: &[PoolConfig { address: Address { configured_role: Role::Replica, - ..Address::new_test() + ..Default::default() }, config: Config::default(), }], diff --git a/pgdog/src/backend/pool/guard.rs b/pgdog/src/backend/pool/guard.rs index 03d27cf76..5833187ea 100644 --- a/pgdog/src/backend/pool/guard.rs +++ b/pgdog/src/backend/pool/guard.rs @@ -349,7 +349,7 @@ mod test { }; let pool = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); pool.launch(); diff --git a/pgdog/src/backend/pool/lb/test.rs b/pgdog/src/backend/pool/lb/test.rs index 50fd6c1b4..41a292cf3 100644 --- a/pgdog/src/backend/pool/lb/test.rs +++ b/pgdog/src/backend/pool/lb/test.rs @@ -2,7 +2,7 @@ use std::collections::HashSet; use std::time::Duration; use tokio::time::sleep; -use crate::backend::pool::{Address, Config, Error, PoolConfig, Request}; +use crate::backend::pool::{Address, Config, Error, PoolConfig, Request, transport::Transport}; use crate::backend::replication::publisher::Lsn; use crate::config::{LoadBalancingStrategy, Role}; use pgdog_stats::{LsnStats as StatsLsnStats, ReplicaLag}; @@ -1250,7 +1250,7 @@ async fn test_move_conns_to_with_added_replica_matches_by_address() { let new_target_for_existing = lb_new .targets .iter() - .find(|t| t.pool.addr().host == "127.0.0.1") + .find(|t| t.pool.addr().host == Transport::new("127.0.0.1")) .expect("should have target for 127.0.0.1"); assert_eq!(new_target_for_existing.role(), Role::Primary); @@ -1258,7 +1258,7 @@ async fn test_move_conns_to_with_added_replica_matches_by_address() { let new_target_for_added = lb_new .targets .iter() - .find(|t| t.pool.addr().host == "localhost") + .find(|t| t.pool.addr().host == Transport::new("localhost")) .expect("should have target for localhost"); assert_eq!(new_target_for_added.role(), Role::Replica); diff --git a/pgdog/src/backend/pool/lsn_monitor.rs b/pgdog/src/backend/pool/lsn_monitor.rs index 891d14607..8c60f92e1 100644 --- a/pgdog/src/backend/pool/lsn_monitor.rs +++ b/pgdog/src/backend/pool/lsn_monitor.rs @@ -450,7 +450,7 @@ mod test { }; let pool = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); pool.launch(); diff --git a/pgdog/src/backend/pool/mod.rs b/pgdog/src/backend/pool/mod.rs index 99430d448..eec4f229a 100644 --- a/pgdog/src/backend/pool/mod.rs +++ b/pgdog/src/backend/pool/mod.rs @@ -25,6 +25,7 @@ pub mod state; pub mod stats; pub mod taken; pub mod token_cache; +pub mod transport; pub mod waiting; pub use address::Address; diff --git a/pgdog/src/backend/pool/monitor.rs b/pgdog/src/backend/pool/monitor.rs index 7cbe0831d..813786916 100644 --- a/pgdog/src/backend/pool/monitor.rs +++ b/pgdog/src/backend/pool/monitor.rs @@ -509,7 +509,7 @@ impl Monitor { #[cfg(test)] mod test { use crate::backend::pool::test::pool; - use crate::backend::pool::{Address, Config, PoolConfig}; + use crate::backend::pool::{Address, Config, PoolConfig, transport::Transport}; use super::*; diff --git a/pgdog/src/backend/pool/pool_impl.rs b/pgdog/src/backend/pool/pool_impl.rs index 376b3cfd3..6bb4a268d 100644 --- a/pgdog/src/backend/pool/pool_impl.rs +++ b/pgdog/src/backend/pool/pool_impl.rs @@ -89,7 +89,7 @@ impl Pool { #[cfg(test)] pub fn new_test() -> Self { let config = PoolConfig { - address: Address::new_test(), + address: Address::default(), config: Config::default(), }; diff --git a/pgdog/src/backend/pool/shard/mod.rs b/pgdog/src/backend/pool/shard/mod.rs index 1dab2e3af..915cb89ab 100644 --- a/pgdog/src/backend/pool/shard/mod.rs +++ b/pgdog/src/backend/pool/shard/mod.rs @@ -390,14 +390,14 @@ mod test { crate::logger(); let primary = Some(&PoolConfig { - address: Address::new_test(), + address: Address::default(), ..Default::default() }); let replicas = &[PoolConfig { address: Address { configured_role: Role::Replica, - ..Address::new_test() + ..Default::default() }, ..Default::default() }]; @@ -430,12 +430,12 @@ mod test { crate::logger(); let primary = Some(&PoolConfig { - address: Address::new_test(), + address: Address::default(), ..Default::default() }); let replicas = &[PoolConfig { - address: Address::new_test(), + address: Address::default(), ..Default::default() }]; @@ -468,7 +468,7 @@ mod test { let replicas = &[PoolConfig { address: Address { configured_role: Role::Auto, - ..Address::new_test() + ..Default::default() }, config: super::super::Config { inner: pgdog_stats::Config { diff --git a/pgdog/src/backend/pool/shard/monitor.rs b/pgdog/src/backend/pool/shard/monitor.rs index af3df05e0..826d3bd5e 100644 --- a/pgdog/src/backend/pool/shard/monitor.rs +++ b/pgdog/src/backend/pool/shard/monitor.rs @@ -250,13 +250,13 @@ mod test { #[test] fn test_update_replica_lag_assigns_primary_minus_replica_to_replica_pool() { let primary = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config: Config::default(), }); let replica = Pool::new(&PoolConfig { address: Address { configured_role: Role::Replica, - ..Address::new_test() + ..Default::default() }, config: Config::default(), }); @@ -282,10 +282,10 @@ mod test { async fn test_monitor_updates_roles_on_failover() { crate::logger(); - let primary = Some(&pool_config(Address::new_test())); + let primary = Some(&pool_config(Address::default())); let replicas = [pool_config(Address { configured_role: Role::Auto, - ..Address::new_test() + ..Default::default() })]; let shard = Shard::new(ShardConfig { diff --git a/pgdog/src/backend/pool/shard/role_detector.rs b/pgdog/src/backend/pool/shard/role_detector.rs index 6fcc8eadc..9b3a8d7a2 100644 --- a/pgdog/src/backend/pool/shard/role_detector.rs +++ b/pgdog/src/backend/pool/shard/role_detector.rs @@ -44,7 +44,7 @@ mod test { use crate::backend::databases::User; use crate::backend::pool::lsn_monitor::LsnStats; - use crate::backend::pool::{Address, Config, PoolConfig}; + use crate::backend::pool::{Address, Config, PoolConfig, transport::Transport}; use crate::backend::replication::publisher::Lsn; use crate::config::{ReadWriteSplit, Role}; use pgdog_stats::LsnStats as StatsLsnStats; diff --git a/pgdog/src/backend/pool/test/mod.rs b/pgdog/src/backend/pool/test/mod.rs index 0d213ded9..d69b1f9e0 100644 --- a/pgdog/src/backend/pool/test/mod.rs +++ b/pgdog/src/backend/pool/test/mod.rs @@ -13,6 +13,7 @@ use tokio_util::task::TaskTracker; use crate::backend::ConnectReason; use crate::backend::pool::token_cache::TokenCache; +use crate::backend::pool::transport::Transport; use crate::net::ProtocolMessage; use crate::net::{Parse, Protocol, Query, Sync}; use crate::state::State; @@ -624,7 +625,7 @@ async fn test_checkout_timeout() { }; let pool = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); pool.launch(); @@ -721,13 +722,13 @@ async fn test_move_conns_all_idle() { }; let source = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); source.launch(); let destination = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); @@ -768,13 +769,13 @@ async fn test_move_conns_all_checked_out() { }; let source = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); source.launch(); let destination = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); @@ -823,13 +824,13 @@ async fn test_move_conns_destination_serves_after_launch() { }; let source = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); source.launch(); let destination = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); @@ -1004,7 +1005,7 @@ async fn test_lsn_monitor() { }; let pool = Pool::new(&PoolConfig { - address: Address::new_test(), + address: Address::default(), config, }); diff --git a/pgdog/src/backend/pool/token_cache.rs b/pgdog/src/backend/pool/token_cache.rs index f751c12f3..309c2f745 100644 --- a/pgdog/src/backend/pool/token_cache.rs +++ b/pgdog/src/backend/pool/token_cache.rs @@ -76,7 +76,7 @@ impl From<&Address> for CacheKey { fn from(addr: &Address) -> Self { Self { user: addr.user.clone(), - host: addr.host.clone(), + host: addr.host.to_string(), port: addr.port, } } @@ -280,6 +280,7 @@ impl TokenCache { #[cfg(test)] mod tests { use super::*; + use crate::backend::pool::transport::Transport; use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; diff --git a/pgdog/src/backend/pool/transport.rs b/pgdog/src/backend/pool/transport.rs new file mode 100644 index 000000000..31c683e3b --- /dev/null +++ b/pgdog/src/backend/pool/transport.rs @@ -0,0 +1,71 @@ +use serde::Deserialize; +use serde::Serialize; +use std::fmt::Display; +use std::path::Path; +use std::path::PathBuf; +use thiserror::Error; + +/// Transport enum +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)] +pub enum Transport { + TCP(String), + Unix(PathBuf), +} + +#[derive(Debug, Error, Clone, PartialEq, Eq)] +pub enum TransportError { + #[error("Expected a TCP host but address is a unix socket directory {0}")] + ExpectedTCP(PathBuf), + + #[error("Expected Unix socket directory but address is a TCP host {0}")] + ExpectedUnix(String), +} + +pub fn unix_socket_path(dir: &Path, port: &u16) -> PathBuf { + dir.join(format!(".s.PGSQL.{}", port)) +} + +impl Transport { + pub fn new(value: &str) -> Self { + if value.starts_with('/') { + Transport::Unix(value.into()) + } else { + Transport::TCP(value.to_string()) + } + } + + pub fn tcp(&self) -> Result<&str, TransportError> { + match self { + Transport::TCP(host) => Ok(host), + Transport::Unix(path) => Err(TransportError::ExpectedTCP(path.clone())), + } + } + + pub fn unix(&self) -> Result<&Path, TransportError> { + match self { + Transport::TCP(host) => Err(TransportError::ExpectedUnix(host.clone())), + Transport::Unix(path_buf) => Ok(path_buf), + } + } +} + +impl Display for Transport { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Transport::TCP(addr) => write!(f, "{}", addr), + Transport::Unix(path_buf) => write!(f, "{}", path_buf.display()), + } + } +} + +impl From<&str> for Transport { + fn from(value: &str) -> Self { + Transport::new(value) + } +} + +impl From for Transport { + fn from(value: String) -> Self { + Transport::new(&value) + } +} diff --git a/pgdog/src/backend/pool/waiting.rs b/pgdog/src/backend/pool/waiting.rs index 8b02aad2e..6480e1bc4 100644 --- a/pgdog/src/backend/pool/waiting.rs +++ b/pgdog/src/backend/pool/waiting.rs @@ -91,6 +91,7 @@ pub(super) struct Waiter { mod tests { use super::*; use crate::backend::pool::Pool; + use crate::backend::pool::transport::Transport; use crate::net::messages::FrontendPid; use tokio::time::{Duration, sleep, timeout}; @@ -157,7 +158,7 @@ mod tests { database_name: "pgdog".into(), user: "pgdog".into(), passwords: vec!["pgdog".into()], - ..Default::default() + ..crate::backend::pool::Address::default() }, config, }); diff --git a/pgdog/src/backend/schema/sync/pg_dump.rs b/pgdog/src/backend/schema/sync/pg_dump.rs index f1a44a97e..a4808b4f7 100644 --- a/pgdog/src/backend/schema/sync/pg_dump.rs +++ b/pgdog/src/backend/schema/sync/pg_dump.rs @@ -186,7 +186,7 @@ fn build_pg_dump_command( command .arg("--schema-only") .arg("-h") - .arg(&addr.host) + .arg(addr.host.to_string()) .arg("-p") .arg(addr.port.to_string()) .arg("-U") @@ -1886,7 +1886,7 @@ mod test { #[test] fn test_build_pg_dump_command_sets_password_env() { - let addr = backend::pool::Address::new_test(); + let addr = backend::pool::Address::default(); let command = build_pg_dump_command("pg_dump", &addr, "secret"); let env = command @@ -1908,7 +1908,7 @@ mod test { #[test] fn test_build_pg_dump_command_sets_tls_for_rds_iam() { - let mut addr = backend::pool::Address::new_test(); + let mut addr = backend::pool::Address::default(); addr.server_auth = ServerAuth::RdsIam; let command = build_pg_dump_command("pg_dump", &addr, "token"); @@ -1923,7 +1923,7 @@ mod test { #[test] fn test_build_pg_dump_command_sets_tls_for_azure_workload_identity() { - let mut addr = backend::pool::Address::new_test(); + let mut addr = backend::pool::Address::default(); addr.server_auth = ServerAuth::AzureWorkloadIdentity; let command = build_pg_dump_command("pg_dump", &addr, "token"); diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 96d9d0e7f..e4e866fda 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -6,7 +6,7 @@ use bytes::{BufMut, BytesMut}; use rustls_pki_types::ServerName; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, - net::TcpStream, + net::{TcpStream, UnixStream}, spawn, time::Instant, }; @@ -18,7 +18,10 @@ use super::{ }; use crate::{ auth::{md5, scram::Client}, - backend::pool::stats::MemoryStats, + backend::pool::{ + stats::MemoryStats, + transport::{Transport, unix_socket_path}, + }, config::AuthType, frontend::ClientRequest, net::{ @@ -238,22 +241,34 @@ impl Server { oids: Arc, ) -> Result { debug!("=> {}", addr); - let stream = TcpStream::connect(addr.addr().await?).await?; let config = config(); - - if let Err(err) = tweak(&stream, &config.config.tcp) { - warn!( - "keepalive settings ({}) are not supported on this system, ignoring, error: {} [{}]", - config.config.tcp, err, addr, - ); - } - - let mut stream = Stream::plain(stream, config.config.memory.net_buffer); + let mut stream = match &addr.host { + Transport::TCP(_) => { + let socket = addr.addr().await?; + let tcp = TcpStream::connect(socket).await?; + + if let Err(err) = tweak(&tcp, &config.config.tcp) { + warn!( + "keepalive settings ({}) are not supported on this system, ignoring, error: {} [{}]", + config.config.tcp, err, addr, + ); + } + Stream::plain(tcp, config.config.memory.net_buffer) + } + Transport::Unix(dir) => { + let path = unix_socket_path(dir, &addr.port); + debug!("connecting to Unix socket {}", path.display()); + Stream::unix( + UnixStream::connect(&path).await?, + config.config.memory.net_buffer, + ) + } + }; let tls_mode = config.config.general.tls_verify; - // Only attempt TLS if not in Disabled mode - if tls_mode != TlsVerifyMode::Disabled { + // Only attempt TLS if not in Disabled mode and its not connecting to a unix socket + if tls_mode != TlsVerifyMode::Disabled && matches!(addr.host, Transport::TCP(_)) { debug!( "requesting TLS connection with verify mode: {:?} [{}]", tls_mode, addr, @@ -278,7 +293,8 @@ impl Server { )?; let plain = stream.take()?; - let server_name = ServerName::try_from(addr.host.clone())?; + let host = addr.host.tcp()?.to_owned(); + let server_name = ServerName::try_from(host)?; debug!("connecting with TLS to server name: {:?}", server_name); match connector.connect(server_name.clone(), plain).await { @@ -450,7 +466,18 @@ impl Server { /// Request query cancellation for the given backend server identifier. pub async fn cancel(addr: &Address, id: BackendKeyData) -> Result<(), Error> { - let mut stream = TcpStream::connect(addr.addr().await?).await?; + let mut stream = match &addr.host { + Transport::TCP(_) => { + let tcp = TcpStream::connect(addr.addr().await?).await?; + Stream::plain(tcp, config().config.memory.net_buffer) + } + Transport::Unix(dir) => { + let path = unix_socket_path(dir, &addr.port); + let unix = UnixStream::connect(&path).await?; + Stream::unix(unix, config().config.memory.net_buffer) + } + }; + stream.write_all(&Startup::Cancel { id }.to_bytes()).await?; stream.flush().await?; @@ -1316,7 +1343,7 @@ pub mod test { use bytes::{BufMut, BytesMut}; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, - net::TcpListener, + net::{TcpListener, UnixListener}, }; use crate::{ @@ -1391,7 +1418,7 @@ pub mod test { pub(crate) async fn test_server() -> Server { Server::connect( - &Address::new_test(), + &Address::default(), ServerOptions::default(), ConnectReason::Other, Default::default(), @@ -1407,7 +1434,7 @@ pub mod test { Server::connect( &Address { database_name: "pgdog1".into(), - ..Address::new_test() + ..Default::default() }, ServerOptions::default(), ConnectReason::Other, @@ -1419,7 +1446,7 @@ pub mod test { pub async fn test_replication_server() -> Server { Server::connect( - &Address::new_test(), + &Address::default(), ServerOptions::new_replication(), ConnectReason::Replication, Default::default(), @@ -1512,7 +1539,7 @@ pub mod test { } }); - let mut addr = Address::new_test(); + let mut addr = Address::default(); addr.port = port; addr.server_auth = crate::config::ServerAuth::RdsIam; addr.server_iam_region = Some("us-east-1".into()); @@ -1579,7 +1606,7 @@ pub mod test { } }); - let mut addr = Address::new_test(); + let mut addr = Address::default(); addr.port = port; addr.server_auth = crate::config::ServerAuth::AzureWorkloadIdentity; addr.passwords = vec!["wrong-password".into()]; @@ -1603,6 +1630,65 @@ pub mod test { server_task.await.unwrap(); } + #[tokio::test] + async fn test_connect_over_unix_socket() { + let dir = std::env::temp_dir().join(format!("pgdog-unix-test-{}", std::process::id())); + std::fs::create_dir_all(&dir).unwrap(); + let port = 6543; + let socket_path = dir.join(format!(".s.PGSQL.{}", port)); + + // Bind the listener at the path pgdog derives from the socket + // directory and port. If the derivation is wrong, connect fails. + let listener = UnixListener::bind(&socket_path).unwrap(); + + let server_task = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + + // TLS is skipped for Unix sockets: the first packet must be the + // Startup message, never an SSLRequest. + let startup = Startup::from_stream(&mut socket).await.unwrap(); + assert!( + matches!(startup, Startup::Startup { .. }), + "expected Startup, got {:?}", + startup + ); + + // Peer/trust auth replies: no password is ever exchanged. + socket + .write_all(&Authentication::Ok.to_bytes()) + .await + .unwrap(); + socket + .write_all(&BackendKeyData::random_legacy().to_bytes()) + .await + .unwrap(); + socket + .write_all(&ReadyForQuery::idle().to_bytes()) + .await + .unwrap(); + }); + + let addr = Address { + host: Transport::new(dir.to_str().unwrap()), + port, + ..Address::default() + }; + + let server = Server::connect( + &addr, + ServerOptions::default(), + ConnectReason::Other, + Default::default(), + ) + .await + .expect("connect over unix socket"); + + drop(server); + server_task.await.unwrap(); + + std::fs::remove_dir_all(&dir).unwrap(); + } + #[tokio::test] async fn test_simple_query() { let mut server = test_server().await; diff --git a/pgdog/src/net/stream.rs b/pgdog/src/net/stream.rs index 10ee0d3e9..23c41f69d 100644 --- a/pgdog/src/net/stream.rs +++ b/pgdog/src/net/stream.rs @@ -3,18 +3,38 @@ use bytes::{BufMut, BytesMut}; use futures::FutureExt; use pin_project::pin_project; -use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufStream, ReadBuf}; -use tokio::net::TcpStream; +use tokio::io::{self, AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufStream, ReadBuf}; +use tokio::net::{TcpStream, UnixStream, unix}; use tracing::trace; use std::io::{Error, ErrorKind}; use std::net::SocketAddr; -use std::ops::Deref; +use std::os::fd::AsRawFd; +use std::path::PathBuf; use std::pin::Pin; use std::task::Context; use super::messages::{ErrorResponse, Message, Protocol, ReadyForQuery}; +fn unix_peek(stream: &UnixStream, buffer: &mut [u8]) -> Option> { + let n = unsafe { + libc::recv( + stream.as_raw_fd(), + buffer.as_mut_ptr().cast(), + buffer.len(), + libc::MSG_PEEK, + ) + }; + + if n >= 0 { + Some(Ok(n as usize)) + } else if io::Error::last_os_error().kind() == io::ErrorKind::WouldBlock { + None + } else { + Some(Err(io::Error::last_os_error())) + } +} + /// Inner stream types. #[pin_project(project = StreamInnerProjection)] #[derive(Debug)] @@ -23,6 +43,7 @@ enum StreamInner { Plain(#[pin] BufStream), Tls(#[pin] BufStream>), DevNull, + UnixSockets(#[pin] BufStream), } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -54,6 +75,7 @@ impl AsyncRead for Stream { match project.inner.project() { StreamInnerProjection::Plain(stream) => stream.poll_read(cx, buf), StreamInnerProjection::Tls(stream) => stream.poll_read(cx, buf), + StreamInnerProjection::UnixSockets(stream) => stream.poll_read(cx, buf), StreamInnerProjection::DevNull => std::task::Poll::Ready(Ok(())), } } @@ -70,6 +92,7 @@ impl AsyncWrite for Stream { StreamInnerProjection::Plain(stream) => stream.poll_write(cx, buf), StreamInnerProjection::Tls(stream) => stream.poll_write(cx, buf), StreamInnerProjection::DevNull => std::task::Poll::Ready(Ok(buf.len())), + StreamInnerProjection::UnixSockets(stream) => stream.poll_write(cx, buf), } } @@ -82,6 +105,7 @@ impl AsyncWrite for Stream { StreamInnerProjection::Plain(stream) => stream.poll_flush(cx), StreamInnerProjection::Tls(stream) => stream.poll_flush(cx), StreamInnerProjection::DevNull => std::task::Poll::Ready(Ok(())), + StreamInnerProjection::UnixSockets(stream) => stream.poll_flush(cx), } } @@ -93,6 +117,7 @@ impl AsyncWrite for Stream { match project.inner.project() { StreamInnerProjection::Plain(stream) => stream.poll_shutdown(cx), StreamInnerProjection::Tls(stream) => stream.poll_shutdown(cx), + StreamInnerProjection::UnixSockets(stream) => stream.poll_shutdown(cx), StreamInnerProjection::DevNull => std::task::Poll::Ready(Ok(())), } } @@ -115,6 +140,17 @@ impl Stream { } } + /// Wrap a unix socket stream. + pub fn unix(stream: UnixStream, capacity: usize) -> Self { + Self { + inner: StreamInner::UnixSockets(BufStream::with_capacity(capacity, capacity, stream)), + io_in_progress: false, + capacity, + tls_identity: None, + tls_client_certificate: false, + } + } + /// Wrap an encrypted TCP stream. pub fn tls( stream: tokio_rustls::TlsStream, @@ -158,13 +194,14 @@ impl Stream { matches!(self.inner, StreamInner::Tls(_)) } - /// Get peer address if any. We're not using UNIX sockets (yet) + /// Get peer address/unix socket if any. /// so the peer address should always be available. pub fn peer_addr(&self) -> PeerAddr { match &self.inner { - StreamInner::Plain(stream) => stream.get_ref().peer_addr().ok().into(), - StreamInner::Tls(stream) => stream.get_ref().get_ref().0.peer_addr().ok().into(), - StreamInner::DevNull => PeerAddr { addr: None }, + StreamInner::Plain(stream) => stream.get_ref().peer_addr().into(), + StreamInner::Tls(stream) => stream.get_ref().get_ref().0.peer_addr().into(), + StreamInner::DevNull => PeerAddr::Empty, + StreamInner::UnixSockets(stream) => stream.get_ref().peer_addr().into(), } } @@ -175,6 +212,10 @@ impl Stream { StreamInner::Plain(plain) => eof(plain.get_mut().peek(&mut buf).await)?, StreamInner::Tls(tls) => eof(tls.get_mut().get_mut().0.peek(&mut buf).await)?, StreamInner::DevNull => 0, + StreamInner::UnixSockets(stream) => match unix_peek(stream.get_ref(), &mut buf) { + None => 0, + Some(res) => eof(res)?, + }, }; Ok(()) @@ -186,6 +227,7 @@ impl Stream { StreamInner::Plain(plain) => plain.get_mut().peek(&mut buf).now_or_never(), StreamInner::Tls(tls) => tls.get_mut().get_mut().0.peek(&mut buf).now_or_never(), StreamInner::DevNull => return Liveness::Clean, + StreamInner::UnixSockets(stream) => unix_peek(stream.get_ref(), &mut buf), }; match peeked { @@ -216,6 +258,7 @@ impl Stream { StreamInner::Plain(stream) => eof(stream.write_all(&bytes).await)?, StreamInner::Tls(stream) => eof(stream.write_all(&bytes).await)?, StreamInner::DevNull => (), + StreamInner::UnixSockets(stream) => eof(stream.write_all(&bytes).await)?, } #[cfg(debug_assertions)] @@ -365,30 +408,44 @@ pub fn eof(result: std::io::Result) -> Result { /// Wrapper around SocketAddr /// to make it easier to debug. -pub struct PeerAddr { - addr: Option, +pub enum PeerAddr { + /// TCP peer: Ip and Port + TCP(SocketAddr), + /// Unix Socket file path + Unix(PathBuf), + /// No Identifiable Peer + Empty, } -impl Deref for PeerAddr { - type Target = Option; - - fn deref(&self) -> &Self::Target { - &self.addr +impl From> for PeerAddr { + fn from(addr: io::Result) -> Self { + match addr { + Ok(addr) => PeerAddr::TCP(addr), + Err(_) => PeerAddr::Empty, + } } } -impl From> for PeerAddr { - fn from(value: Option) -> Self { - Self { addr: value } +impl From> for PeerAddr { + fn from(addr: io::Result) -> Self { + match addr { + Ok(addr) => { + let Some(path) = addr.as_pathname() else { + return PeerAddr::Empty; + }; + PeerAddr::Unix(path.into()) + } + Err(_) => PeerAddr::Empty, + } } } impl std::fmt::Debug for PeerAddr { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - if let Some(addr) = &self.addr { - write!(f, "[{}]", addr) - } else { - write!(f, "") + match self { + PeerAddr::TCP(addr) => write!(f, "[{}]", addr), + PeerAddr::Unix(file_path) => write!(f, "{}", file_path.display()), + PeerAddr::Empty => write!(f, "No address"), } } } @@ -398,7 +455,10 @@ mod tests { use std::time::Duration; use super::*; - use tokio::net::TcpListener; + use tokio::{ + io::AsyncWriteExt, + net::{TcpListener, UnixListener, UnixStream}, + }; #[tokio::test] async fn test_io_in_progress_initially_false() { @@ -516,4 +576,213 @@ mod tests { assert_eq!(stream.liveness(), Liveness::Clean); } + + // ── Unix domain socket streams ────────────────────────────────────────── + + /// Bind a Unix socket pair and wrap the *connecting* end in a [`Stream`] + /// (pgdog is the connecting client; the listener is the Postgres side). + async fn unix_pair(name: &str) -> (Stream, UnixStream, std::path::PathBuf) { + let dir = + std::env::temp_dir().join(format!("pgdog-stream-{}-{}", name, std::process::id())); + std::fs::create_dir_all(&dir).unwrap(); + let path = dir.join("test.sock"); + let listener = UnixListener::bind(&path).unwrap(); + let server = tokio::spawn(async move { + let (accepted, _) = listener.accept().await.unwrap(); + accepted + }); + let client = UnixStream::connect(&path).await.unwrap(); + let stream = Stream::unix(client, 4096); + (stream, server.await.unwrap(), dir) + } + + #[tokio::test] + async fn test_unix_stream_peer_addr() { + let (stream, _peer, dir) = unix_pair("peer").await; + + let peer_addr = stream.peer_addr(); + let expected = dir.join("test.sock"); + match &peer_addr { + PeerAddr::Unix(path) => assert_eq!(path, &expected), + other => panic!("expected PeerAddr::Unix({:?}), got {:?}", expected, other), + } + + drop(stream); + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_unix_stream_liveness() { + let (mut stream, mut peer, dir) = unix_pair("liveness").await; + + // Idle: the socket is open but has no data. + assert_eq!(stream.liveness(), Liveness::Clean); + + // Unsolicited data is visible to liveness without consuming it. + peer.write_all(b"x").await.unwrap(); + peer.flush().await.unwrap(); + tokio::task::yield_now().await; + assert_eq!(stream.liveness(), Liveness::DataPending); + + // The peek must not consume the byte. + let mut buf = [0u8; 1]; + stream.read_exact(&mut buf).await.unwrap(); + assert_eq!(&buf, b"x"); + + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_unix_stream_shutdown() { + let (mut stream, _peer, dir) = unix_pair("shutdown").await; + + // Graceful shutdown polls the UnixSockets shutdown arm. + stream.shutdown().await.unwrap(); + + drop(stream); + std::fs::remove_dir_all(&dir).unwrap(); + } + + // ── unix_peek edge cases ──────────────────────────────────────────────── + + /// Raw Unix socket pair (connecting end, accepted end) for direct + /// [`unix_peek`] tests. + async fn unix_stream_pair(name: &str) -> (UnixStream, UnixStream, std::path::PathBuf) { + let dir = std::env::temp_dir().join(format!("pgdog-peek-{}-{}", name, std::process::id())); + std::fs::create_dir_all(&dir).unwrap(); + let path = dir.join("test.sock"); + let listener = UnixListener::bind(&path).unwrap(); + let server = tokio::spawn(async move { + let (accepted, _) = listener.accept().await.unwrap(); + accepted + }); + let client = UnixStream::connect(&path).await.unwrap(); + (client, server.await.unwrap(), dir) + } + + #[tokio::test] + async fn test_unix_peek_no_data_is_none() { + let (client, _server, dir) = unix_stream_pair("no-data").await; + + // Idle non-blocking socket: recv returns WouldBlock → None. + let mut buf = [0u8; 1]; + assert!(unix_peek(&client, &mut buf).is_none()); + + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_unix_peek_reports_data_without_consuming() { + let (mut client, mut server, dir) = unix_stream_pair("peek").await; + + server.write_all(b"hello").await.unwrap(); + server.flush().await.unwrap(); + tokio::task::yield_now().await; + + let mut buf = [0u8; 1]; + assert!(matches!(unix_peek(&client, &mut buf), Some(Ok(1)))); + + // MSG_PEEK must not consume: a full read still gets everything. + let mut out = [0u8; 5]; + client.read_exact(&mut out).await.unwrap(); + assert_eq!(&out, b"hello"); + + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_unix_peek_limited_by_buffer_size() { + let (mut client, mut server, dir) = unix_stream_pair("limit").await; + + server.write_all(&[0u8; 100]).await.unwrap(); + server.flush().await.unwrap(); + tokio::task::yield_now().await; + + // Peek returns min(available, buffer length), not the full payload. + let mut buf = [0u8; 10]; + assert!(matches!(unix_peek(&client, &mut buf), Some(Ok(10)))); + + // And the remaining 90 bytes are still there. + let mut out = vec![0u8; 100]; + client.read_exact(&mut out).await.unwrap(); + + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_unix_peek_eof_after_peer_close() { + let (client, server, dir) = unix_stream_pair("eof").await; + drop(server); + tokio::task::yield_now().await; + + // Peer closed: recv returns 0 → Some(Ok(0)) → liveness() = Closed. + let mut buf = [0u8; 1]; + assert!(matches!(unix_peek(&client, &mut buf), Some(Ok(0)))); + + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_unix_peek_error_on_bad_fd() { + let (client, _server, dir) = unix_stream_pair("bad-fd").await; + + // Close the fd behind the stream's back: recv then fails with EBADF, + // exercising the Err branch of unix_peek. + unsafe { libc::close(client.as_raw_fd()) }; + + let mut buf = [0u8; 1]; + let result = unix_peek(&client, &mut buf); + assert!( + matches!(result, Some(Err(ref e)) if e.raw_os_error() == Some(libc::EBADF)), + "expected Some(Err(EBADF)), got {:?}", + result + ); + + // The fd is already closed; forget the stream so its Drop doesn't + // close a possibly-reused fd number. + std::mem::forget(client); + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[tokio::test] + async fn test_peer_addr_conversions() { + let addr: SocketAddr = "127.0.0.1:5432".parse().unwrap(); + assert!(matches!(PeerAddr::from(Ok(addr)), PeerAddr::TCP(_))); + + let tcp_err: io::Result = Err(io::Error::new(io::ErrorKind::Other, "no peer")); + assert!(matches!(PeerAddr::from(tcp_err), PeerAddr::Empty)); + + let unix_err: io::Result = + Err(io::Error::new(io::ErrorKind::Other, "no peer")); + assert!(matches!(PeerAddr::from(unix_err), PeerAddr::Empty)); + + // A Unix socket whose peer has no pathname (e.g. the accepted end of + // an unbound connecting client) maps to Empty, not Unix. + let dir = std::env::temp_dir().join(format!("pgdog-stream-conv-{}", std::process::id())); + std::fs::create_dir_all(&dir).unwrap(); + let path = dir.join("test.sock"); + let listener = UnixListener::bind(&path).unwrap(); + let server = tokio::spawn(async move { + let (accepted, _) = listener.accept().await.unwrap(); + accepted.peer_addr().unwrap() + }); + let _client = UnixStream::connect(&path).await.unwrap(); + let unnamed = server.await.unwrap(); + assert!(unnamed.as_pathname().is_none()); + assert!(matches!(PeerAddr::from(Ok(unnamed)), PeerAddr::Empty)); + std::fs::remove_dir_all(&dir).unwrap(); + } + + #[test] + fn test_peer_addr_debug() { + let addr: SocketAddr = "127.0.0.1:5432".parse().unwrap(); + assert_eq!(format!("{:?}", PeerAddr::TCP(addr)), "[127.0.0.1:5432]"); + assert_eq!(format!("{:?}", PeerAddr::Empty), "No address"); + + let path = PathBuf::from("/var/run/postgresql/.s.PGSQL.5432"); + assert_eq!( + format!("{:?}", PeerAddr::Unix(path)), + "/var/run/postgresql/.s.PGSQL.5432" + ); + } } diff --git a/pgdog/src/stats/pools.rs b/pgdog/src/stats/pools.rs index 8f8dbee44..3082347b1 100644 --- a/pgdog/src/stats/pools.rs +++ b/pgdog/src/stats/pools.rs @@ -103,7 +103,7 @@ impl Pools { let labels = vec![ ("user".into(), user.user.clone()), ("database".into(), user.database.clone()), - ("host".into(), pool.addr().host.clone()), + ("host".into(), pool.addr().host.to_string()), ("port".into(), pool.addr().port.to_string()), ("shard".into(), shard_num.to_string()), ("role".into(), role.to_string()),