From 5f46f3cdd273956e591811ddb5250ce73e4628f8 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Tue, 6 Oct 2026 18:41:28 -0700 Subject: [PATCH] feat: rds iam passthrough auth --- integration/complex/rds_iam/dev.sh | 14 ++ integration/complex/rds_iam/pgdog.toml | 11 ++ integration/complex/rds_iam/requirements.txt | 3 + integration/complex/rds_iam/test_rds_iam.py | 54 ++++++++ integration/complex/rds_iam/users.toml | 4 + pgdog-config/src/auth.rs | 9 +- pgdog-config/src/users.rs | 7 +- pgdog/src/auth/mod.rs | 1 + pgdog/src/auth/token_cache.rs | 108 +++++++++++++++ pgdog/src/backend/auth/rds_iam.rs | 9 +- pgdog/src/backend/connect_reason.rs | 1 + pgdog/src/backend/databases.rs | 5 +- pgdog/src/backend/disconnect_reason.rs | 14 +- pgdog/src/backend/pool/address.rs | 13 ++ pgdog/src/backend/pool/cluster.rs | 12 ++ pgdog/src/backend/pool/connection_creation.rs | 131 ++++++++++++++++++ pgdog/src/backend/pool/mod.rs | 2 + pgdog/src/backend/pool/monitor.rs | 93 ++----------- pgdog/src/backend/pool/password.rs | 6 +- pgdog/src/backend/pool/pool_impl.rs | 27 +++- pgdog/src/backend/server.rs | 7 +- pgdog/src/frontend/client/mod.rs | 25 +++- 22 files changed, 445 insertions(+), 111 deletions(-) create mode 100755 integration/complex/rds_iam/dev.sh create mode 100644 integration/complex/rds_iam/pgdog.toml create mode 100644 integration/complex/rds_iam/requirements.txt create mode 100644 integration/complex/rds_iam/test_rds_iam.py create mode 100644 integration/complex/rds_iam/users.toml create mode 100644 pgdog/src/auth/token_cache.rs create mode 100644 pgdog/src/backend/pool/connection_creation.rs diff --git a/integration/complex/rds_iam/dev.sh b/integration/complex/rds_iam/dev.sh new file mode 100755 index 000000000..d1b431887 --- /dev/null +++ b/integration/complex/rds_iam/dev.sh @@ -0,0 +1,14 @@ +#!/bin/bash +set -euo pipefail +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd ) + +cd "${SCRIPT_DIR}" + +if [[ ! -x venv/bin/python ]]; then + python3 -m venv venv +fi + +venv/bin/python -m pip install -r requirements.txt + +# Requires PgDog on 127.0.0.1:6432 and configured AWS credentials. +venv/bin/python -m pytest test_rds_iam.py -v --tb=short "$@" diff --git a/integration/complex/rds_iam/pgdog.toml b/integration/complex/rds_iam/pgdog.toml new file mode 100644 index 000000000..2eacc53f4 --- /dev/null +++ b/integration/complex/rds_iam/pgdog.toml @@ -0,0 +1,11 @@ +[general] +host = "127.0.0.1" +port = 6432 +auth_type = "external_token" + +[[databases]] +name = "postgres" +host = "staging-iam.cluster-c5icciqq4b0q.us-west-2.rds.amazonaws.com" + +[admin] +password = "pgdog" diff --git a/integration/complex/rds_iam/requirements.txt b/integration/complex/rds_iam/requirements.txt new file mode 100644 index 000000000..7c2761bbf --- /dev/null +++ b/integration/complex/rds_iam/requirements.txt @@ -0,0 +1,3 @@ +boto3 +psycopg==3.2.6 +pytest==8.3.5 diff --git a/integration/complex/rds_iam/test_rds_iam.py b/integration/complex/rds_iam/test_rds_iam.py new file mode 100644 index 000000000..99b33e0c1 --- /dev/null +++ b/integration/complex/rds_iam/test_rds_iam.py @@ -0,0 +1,54 @@ +from urllib.parse import parse_qs, urlsplit + +import boto3 +import psycopg +import pytest +from botocore.exceptions import NoCredentialsError + + +@pytest.fixture(scope="module") +def rds_iam_token(): + """Generate a staging-iam token using the default AWS credential chain.""" + hostname = "staging-iam.cluster-c5icciqq4b0q.us-west-2.rds.amazonaws.com" + client = boto3.client("rds", region_name="us-west-2") + + try: + return client.generate_db_auth_token( + DBHostname=hostname, + Port=5432, + DBUsername="postgres", + Region="us-west-2", + ) + except NoCredentialsError: + pytest.skip("AWS credentials are required to generate an RDS IAM token") + + +def test_generate_rds_iam_token(rds_iam_token): + url = urlsplit(f"https://{rds_iam_token}") + params = parse_qs(url.query) + assert ( + url.hostname == "staging-iam.cluster-c5icciqq4b0q.us-west-2.rds.amazonaws.com" + ) + assert url.port == 5432 + assert params["Action"] == ["connect"] + assert params["DBUser"] == ["postgres"] + assert params["X-Amz-Expires"] == ["900"] + assert params["X-Amz-Credential"][0].endswith("/us-west-2/rds-db/aws4_request") + assert params["X-Amz-Signature"][0] + + +def test_connect_to_pgdog(rds_iam_token): + """Reuse one IAM token across ten separate client connections.""" + for attempt in range(10): + with psycopg.connect( + host="127.0.0.1", + port=6432, + dbname="postgres", + user="postgres", + password=rds_iam_token, + sslmode="disable", + connect_timeout=10, + options="-c statement_timeout=10000", + ) as conn: + row = conn.execute("SELECT current_user, current_database()").fetchone() + assert row == ("postgres", "postgres"), f"Connection {attempt + 1}" diff --git a/integration/complex/rds_iam/users.toml b/integration/complex/rds_iam/users.toml new file mode 100644 index 000000000..8084039a3 --- /dev/null +++ b/integration/complex/rds_iam/users.toml @@ -0,0 +1,4 @@ +[[users]] +name = "postgres" +database = "postgres" +server_auth = "rds_iam" diff --git a/pgdog-config/src/auth.rs b/pgdog-config/src/auth.rs index 2d938118b..7fad48557 100644 --- a/pgdog-config/src/auth.rs +++ b/pgdog-config/src/auth.rs @@ -37,7 +37,7 @@ impl PassthroughAuth { /// See [authentication](https://docs.pgdog.dev/features/authentication/). /// /// -#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq, JsonSchema)] +#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq, JsonSchema, Copy)] #[serde(rename_all = "snake_case")] pub enum AuthType { /// MD5 password hashing; very quick but not secure. @@ -49,6 +49,8 @@ pub enum AuthType { Trust, /// Plaintext password. Plain, + /// RDS IAM or Azure Workload identity. + ExternalToken, } impl Display for AuthType { @@ -58,6 +60,7 @@ impl Display for AuthType { Self::Scram => write!(f, "scram"), Self::Trust => write!(f, "trust"), Self::Plain => write!(f, "plain"), + Self::ExternalToken => write!(f, "external_token"), } } } @@ -74,6 +77,10 @@ impl AuthType { pub fn trust(&self) -> bool { matches!(self, Self::Trust) } + + pub fn external_token(&self) -> bool { + matches!(self, Self::ExternalToken) + } } impl FromStr for AuthType { diff --git a/pgdog-config/src/users.rs b/pgdog-config/src/users.rs index 86210e32f..7f1852c6b 100644 --- a/pgdog-config/src/users.rs +++ b/pgdog-config/src/users.rs @@ -51,9 +51,12 @@ impl Users { pub fn check(&mut self, config: &Config) { for user in &mut self.users { if user.passwords().is_empty() { - if !config.general.passthrough_auth() && user.identity.is_none() { + if !config.general.passthrough_auth() + && user.identity.is_none() + && !config.general.auth_type.external_token() + { warn!( - r#"user "{}" (database "{}") doesn't have a password, passthrough auth and mTLS are disabled"#, + r#"user "{}" (database "{}") doesn't have a password, while passthrough auth, external token, and mTLS auth are disabled"#, user.name, user.database, ); } diff --git a/pgdog/src/auth/mod.rs b/pgdog/src/auth/mod.rs index 0d5f99b7e..d16c039ed 100644 --- a/pgdog/src/auth/mod.rs +++ b/pgdog/src/auth/mod.rs @@ -4,6 +4,7 @@ pub(crate) mod auth_result; pub(crate) mod error; pub(crate) mod md5; pub(crate) mod scram; +pub(crate) mod token_cache; pub(crate) mod vault; pub(crate) use auth_result::AuthResult; diff --git a/pgdog/src/auth/token_cache.rs b/pgdog/src/auth/token_cache.rs new file mode 100644 index 000000000..22673327a --- /dev/null +++ b/pgdog/src/auth/token_cache.rs @@ -0,0 +1,108 @@ +//! Token cache that validates a token we received +//! from the client is actually valid with whatever mechanism +//! we authenticate to Postgres with. +//! +//! This is effectively passthrough auth for RDS IAM, Azure Workload Identity, etc. + +use moka::sync::Cache; +use once_cell::sync::Lazy; +use rand::seq::{IndexedRandom, IteratorRandom}; +use std::{ + sync::Arc, + time::{Duration, SystemTime}, +}; +use tokio::{sync::Mutex, time::sleep}; + +use crate::backend::pool::Connection; +use crate::tasks::{shutdown_signal, spawn}; + +pub static AUTH_TOKEN_CACHE: Lazy = Lazy::new(|| { + let cache = AuthTokenCache::default(); + let background = cache.clone(); + + spawn("auth token cache", async move { + let shutdown = shutdown_signal(); + loop { + tokio::select! { + _ = sleep(Duration::from_millis(333)) => background.evict(), + _ = shutdown.cancelled() => break, + } + } + }); + + cache +}); + +#[derive(Debug, Clone)] +struct Entry { + expires_at: SystemTime, + lock: Arc>, +} + +#[derive(Clone, Debug)] +pub(crate) struct AuthTokenCache { + tokens: Cache, +} + +impl Default for AuthTokenCache { + fn default() -> Self { + Self { + tokens: Cache::builder() + .max_capacity(1_000) + .support_invalidation_closures() + .build(), + } + } +} + +impl AuthTokenCache { + fn evict(&self) { + let now = SystemTime::now(); + self.tokens + .invalidate_entries_if(move |_, value| value.expires_at <= now) + .expect("invalidation closures are enabled"); + } + + /// Check a token against the external token provider, e.g., RDS. + /// + /// Takes a lock on the token so a connection storm from clients + /// only creates one connection to the actual database. + pub(crate) async fn check( + &self, + user: &str, + database: &str, + token: &str, + ) -> Result { + if self.tokens.contains_key(token) { + return Ok(true); + } + + let lock = { + let entry = self.tokens.entry(token.to_owned()).or_insert(Entry { + expires_at: crate::backend::auth::rds_iam::expires_at(), + lock: Arc::new(Mutex::new(())), + }); + entry.value().lock.clone() + }; + + let _guard = lock.lock().await; + + let conn = Connection::new(user, database, false)?; + let shard = conn + .cluster()? + .shards() + .choose(&mut rand::rng()) + .expect("to have at least one shard"); + let pool = shard + .pool_iter() + .choose(&mut rand::rng()) + .expect("to have at least one pool"); + + if let Ok(()) = pool.validate_token(token).await { + Ok(true) + } else { + self.tokens.remove(token); + Ok(false) + } + } +} diff --git a/pgdog/src/backend/auth/rds_iam.rs b/pgdog/src/backend/auth/rds_iam.rs index 0eca0e494..36dd98d09 100644 --- a/pgdog/src/backend/auth/rds_iam.rs +++ b/pgdog/src/backend/auth/rds_iam.rs @@ -135,9 +135,12 @@ pub(crate) async fn token(addr: Address) -> Result<(String, SystemTime), Error> )) })?; - // RDS IAM tokens are valid for 15 minutes. - let expires_at = SystemTime::now() + Duration::from_secs(900); - Ok((token, expires_at)) + Ok((token, expires_at())) +} + +// RDS IAM tokens are valid for 15 minutes. +pub(crate) fn expires_at() -> SystemTime { + SystemTime::now() + Duration::from_secs(900) } #[cfg(test)] diff --git a/pgdog/src/backend/connect_reason.rs b/pgdog/src/backend/connect_reason.rs index 764ac6473..1058effce 100644 --- a/pgdog/src/backend/connect_reason.rs +++ b/pgdog/src/backend/connect_reason.rs @@ -8,6 +8,7 @@ pub(crate) enum ConnectReason { PubSub, Probe, Healthcheck, + ValidateToken, #[default] Other, } diff --git a/pgdog/src/backend/databases.rs b/pgdog/src/backend/databases.rs index f56e25c5d..35a5d0294 100644 --- a/pgdog/src/backend/databases.rs +++ b/pgdog/src/backend/databases.rs @@ -470,7 +470,10 @@ impl Databases { // Launch all clusters for cluster in self.all().values() { - if cluster.passwords().is_empty() && cluster.identity().is_none() { + if cluster.passwords().is_empty() + && cluster.identity().is_none() + && !cluster.auth_type().external_token() + { warn!( r#"disabling pool for user "{}" and database "{}", password not set"#, cluster.user(), diff --git a/pgdog/src/backend/disconnect_reason.rs b/pgdog/src/backend/disconnect_reason.rs index 52e7b3e32..4a01cff05 100644 --- a/pgdog/src/backend/disconnect_reason.rs +++ b/pgdog/src/backend/disconnect_reason.rs @@ -12,6 +12,7 @@ pub(crate) enum DisconnectReason { Unhealthy, Healthcheck, CredentialsRefresh, + CredentialsCheck, ServerClosed, #[default] Other, @@ -24,14 +25,15 @@ impl Display for DisconnectReason { Self::Old => "max age", Self::Error => "error", Self::Other => "other", - Self::ForceClose => "force close", - Self::Offline => "pool offline", + Self::ForceClose => "force_close", + Self::Offline => "pool_offline", Self::OutOfSync => "out of sync", - Self::ReplicationMode => "in replication mode", + Self::ReplicationMode => "in_replication_mode", Self::Unhealthy => "unhealthy", - Self::Healthcheck => "standalone healthcheck", - Self::CredentialsRefresh => "credentials refresh", - Self::ServerClosed => "server closed", + Self::Healthcheck => "standalone_healthcheck", + Self::CredentialsRefresh => "credentials_refresh", + Self::ServerClosed => "server_closed", + Self::CredentialsCheck => "credentials_check", }; write!(f, "{}", reason) diff --git a/pgdog/src/backend/pool/address.rs b/pgdog/src/backend/pool/address.rs index 654e6b54d..ed3df6cb0 100644 --- a/pgdog/src/backend/pool/address.rs +++ b/pgdog/src/backend/pool/address.rs @@ -193,6 +193,19 @@ impl Address { } } + /// Create an instance of the same address, but swap the password + /// and authentication mechanism so the server is forced to use this password + /// no matter what `server_auth` is configured. + /// + /// This is used for validating RDS IAM tokens we receive from clients. + pub(super) fn into_client_token(self, token: &str) -> Self { + Self { + passwords: vec![Password::new(token, PasswordSource::ClientToken)], + server_auth: ServerAuth::Password, + ..self + } + } + /// Test convention: `new_test()` represents a primary. Tests that need /// a replica do `Address { configured_role: Role::Replica, ..new_test() }`. #[cfg(test)] diff --git a/pgdog/src/backend/pool/cluster.rs b/pgdog/src/backend/pool/cluster.rs index 74c41f40c..c0cc9e6e3 100644 --- a/pgdog/src/backend/pool/cluster.rs +++ b/pgdog/src/backend/pool/cluster.rs @@ -2,6 +2,7 @@ use futures::future::try_join_all; use parking_lot::Mutex; +use pgdog_config::AuthType; use pgdog_config::{ LoadSchema, PreparedStatementsLevel, QueryParser, QueryParserLevel, Rewrite, RewriteMode, users::PasswordKind, @@ -91,6 +92,7 @@ pub(crate) struct Cluster { read_only: bool, failover_signal: ClusterFailoverSignalWatcher, cancellation_token: CancellationToken, + auth_type: AuthType, } /// Bare test clusters carry the same defaults the config would apply, @@ -141,6 +143,7 @@ impl Default for Cluster { read_only: Default::default(), failover_signal: ClusterFailoverSignalWatcher::default(), cancellation_token: Default::default(), + auth_type: AuthType::default(), } } } @@ -229,6 +232,7 @@ pub(crate) struct ClusterConfig<'a> { schema_cache: SchemaCache, canonicalize_oids: bool, read_only: bool, + auth_type: AuthType, } impl<'a> ClusterConfig<'a> { @@ -299,6 +303,7 @@ impl<'a> ClusterConfig<'a> { schema_cache, canonicalize_oids: general.canonicalize_type_information, read_only: user.read_only.unwrap_or(false), + auth_type: general.auth_type, } } } @@ -346,6 +351,7 @@ impl Cluster { schema_cache, canonicalize_oids, read_only, + auth_type, } = config; let identifier = Arc::new(DatabaseUser { @@ -421,6 +427,7 @@ impl Cluster { read_only, failover_signal, cancellation_token: Default::default(), + auth_type, } } @@ -732,6 +739,11 @@ impl Cluster { self.resharding_replication_retry_min_delay } + /// Get auth algorithm used by clients to connect. + pub(crate) fn auth_type(&self) -> &AuthType { + &self.auth_type + } + /// Send a cancellation request for all running queries. pub(crate) async fn cancel_all(&self) -> Result<(), Error> { let pools: Vec<_> = self diff --git a/pgdog/src/backend/pool/connection_creation.rs b/pgdog/src/backend/pool/connection_creation.rs new file mode 100644 index 000000000..c26abb484 --- /dev/null +++ b/pgdog/src/backend/pool/connection_creation.rs @@ -0,0 +1,131 @@ +//! Connection creation primitives. + +use std::{sync::Arc, time::Duration}; +use tokio::time::Instant; + +use crate::util::{safe_sleep, safe_timeout}; +use tracing::error; + +use super::{ + super::{ConnectReason, Oids, Pool, Server, ServerOptions}, + Address, Error, +}; + +pub(super) struct ConnectionArgs<'a> { + address: &'a Address, + timeout: Duration, + attempts: u64, + delay: Duration, + options: ServerOptions, + reason: ConnectReason, + max_age: Duration, + max_age_jitter: Duration, + oids: Arc, + pool: &'a Pool, +} + +impl<'a> ConnectionArgs<'a> { + pub(super) fn from_pool(pool: &'a Pool, reason: ConnectReason) -> Self { + Self { + address: pool.addr(), + timeout: pool.config().connect_timeout, + attempts: pool.config().connect_attempts, + delay: pool.config().connect_attempt_delay, + options: pool.server_options(), + reason, + oids: pool.inner().oids.clone(), + max_age: pool.config().max_age, + max_age_jitter: pool.config().max_age_jitter, + pool, + } + } + + pub(super) fn with_addr(self, addr: &'a Address) -> Self { + Self { + address: addr, + ..self + } + } +} + +pub(super) async fn create(args: ConnectionArgs<'_>) -> Result { + let ConnectionArgs { + address, + timeout, + attempts, + delay, + options, + reason, + oids, + max_age, + max_age_jitter, + pool, + } = args; + + let mut error = Error::ServerError; + let now = Instant::now(); + + for attempt in 0..attempts { + match safe_timeout( + timeout, + Box::pin(Server::connect( + address, + options.clone(), + reason, + oids.clone(), + )), + ) + .await + { + Ok(Ok(mut conn)) => { + conn.stats_mut().set_pool_id(pool.id()); + let elapsed = now.elapsed(); + { + let mut guard = pool.lock(); + guard.stats.counts.connect_count += 1; + guard.stats.counts.connect_time += elapsed; + guard.stats.counts.auth_attempts += conn.password_attempts(); + conn.set_credentials_generation(guard.credentials_generation()); + } + conn.apply_lifetime_jitter(max_age, max_age_jitter); + pool.cache_params(conn.params()); + return Ok(conn); + } + + Ok(Err(err)) => { + // We tried all passwords and they were all wrong. + if err.is_auth() { + pool.lock().stats.counts.auth_attempts += pool.addr().passwords.len(); + } + error!( + "{}error connecting to server: {} [{}]", + if attempt > 0 { + format!("[attempt {}] ", attempt) + } else { + String::new() + }, + err, + pool.addr(), + ); + error = Error::ServerError; + } + + Err(_) => { + error!( + "{}server connection timeout [{}]", + if attempt > 0 { + format!("[attempt {}] ", attempt) + } else { + String::new() + }, + pool.addr(), + ); + error = Error::ConnectTimeout; + } + } + + safe_sleep(delay).await; + } + + Err(error) +} diff --git a/pgdog/src/backend/pool/mod.rs b/pgdog/src/backend/pool/mod.rs index 438f7cf29..23c9d1fcc 100644 --- a/pgdog/src/backend/pool/mod.rs +++ b/pgdog/src/backend/pool/mod.rs @@ -6,6 +6,7 @@ pub(crate) mod cluster; pub(crate) mod cluster_metrics; pub(crate) mod comms; pub(crate) mod connection; +pub(crate) mod connection_creation; pub(crate) mod dns_cache; pub(crate) mod ee; pub(crate) mod error; @@ -50,6 +51,7 @@ pub(crate) use stats::Stats; pub use pgdog_config::pool::PoolConfig as Config; use comms::Comms; +use connection_creation::ConnectionArgs; use inner::Inner; use shard::ShardConfig; use taken::Taken; diff --git a/pgdog/src/backend/pool/monitor.rs b/pgdog/src/backend/pool/monitor.rs index 8aefece2e..be4c0621a 100644 --- a/pgdog/src/backend/pool/monitor.rs +++ b/pgdog/src/backend/pool/monitor.rs @@ -43,18 +43,17 @@ //! The loop exits when the pool shuts down (e.g. on config reload), preventing //! refresh tasks from leaking across reloads. -use std::sync::Arc; use std::time::Duration; use super::{Error, Guard, Healtcheck, Pool, Request}; use crate::backend::auth::{azure_workload_identity, rds_iam, vault}; use crate::backend::pool::inner::ShouldCreate; use crate::backend::pool::token_cache::TokenCache; -use crate::backend::{ConnectReason, DisconnectReason, Server}; +use crate::backend::{ConnectReason, DisconnectReason}; use crate::config::ServerAuth; use crate::tasks; -use crate::util::{safe_interval, safe_sleep, safe_timeout}; +use crate::util::{safe_interval, safe_sleep}; use tokio::select; use tokio::time::Instant; use tracing::{debug, error, info, warn}; @@ -333,7 +332,9 @@ impl Monitor { /// Replenish pool with one new connection. async fn replenish(&self, reason: ConnectReason) -> Result { - match Self::create_connection(&self.pool, reason).await { + let args = super::ConnectionArgs::from_pool(&self.pool, reason); + + match super::connection_creation::create(args).await { Ok(conn) => { let now = Instant::now(); let server = Box::new(conn); @@ -386,7 +387,9 @@ impl Monitor { // Create a new one and close it. info!("creating new healthcheck connection [{}]", pool.addr()); - let mut server = Self::create_connection(pool, ConnectReason::Healthcheck) + let args = super::ConnectionArgs::from_pool(pool, ConnectReason::Healthcheck); + + let mut server = super::connection_creation::create(args) .await .map_err(|_| Error::HealthcheckError)?; @@ -426,86 +429,6 @@ impl Monitor { } } } - - pub(super) async fn create_connection( - pool: &Pool, - reason: ConnectReason, - ) -> Result { - let connect_timeout = pool.config().connect_timeout; - let connect_attempts = pool.config().connect_attempts; - let connect_attempt_delay = pool.config().connect_attempt_delay; - let options = pool.server_options(); - - let mut error = Error::ServerError; - let now = Instant::now(); - - let max_age = pool.config().max_age; - let max_age_jitter = pool.config().max_age_jitter; - - for attempt in 0..connect_attempts { - match safe_timeout( - connect_timeout, - Box::pin(Server::connect( - pool.addr(), - options.clone(), - reason, - Arc::clone(&pool.inner().oids), - )), - ) - .await - { - Ok(Ok(mut conn)) => { - conn.stats_mut().set_pool_id(pool.id()); - let elapsed = now.elapsed(); - { - let mut guard = pool.lock(); - guard.stats.counts.connect_count += 1; - guard.stats.counts.connect_time += elapsed; - guard.stats.counts.auth_attempts += conn.password_attempts(); - conn.set_credentials_generation(guard.credentials_generation()); - } - conn.apply_lifetime_jitter(max_age, max_age_jitter); - pool.cache_params(conn.params()); - return Ok(conn); - } - - Ok(Err(err)) => { - // We tried all passwords and they were all wrong. - if err.is_auth() { - pool.lock().stats.counts.auth_attempts += pool.addr().passwords.len(); - } - error!( - "{}error connecting to server: {} [{}]", - if attempt > 0 { - format!("[attempt {}] ", attempt) - } else { - String::new() - }, - err, - pool.addr(), - ); - error = Error::ServerError; - } - - Err(_) => { - error!( - "{}server connection timeout [{}]", - if attempt > 0 { - format!("[attempt {}] ", attempt) - } else { - String::new() - }, - pool.addr(), - ); - error = Error::ConnectTimeout; - } - } - - safe_sleep(connect_attempt_delay).await; - } - - Err(error) - } } #[cfg(test)] diff --git a/pgdog/src/backend/pool/password.rs b/pgdog/src/backend/pool/password.rs index e0b3a4b8c..950a4f350 100644 --- a/pgdog/src/backend/pool/password.rs +++ b/pgdog/src/backend/pool/password.rs @@ -16,15 +16,17 @@ pub(crate) enum PasswordSource { RdsIam, AzureIdentity, Vault, + ClientToken, } impl Display for PasswordSource { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { Self::Config => write!(f, "config"), - Self::RdsIam => write!(f, "rds iam"), - Self::AzureIdentity => write!(f, "azure workload identity"), + Self::RdsIam => write!(f, "rds_iam"), + Self::AzureIdentity => write!(f, "azure_workload_identity"), Self::Vault => write!(f, "vault"), + Self::ClientToken => write!(f, "client_token"), } } } diff --git a/pgdog/src/backend/pool/pool_impl.rs b/pgdog/src/backend/pool/pool_impl.rs index dd207dd6c..6e9d8b2e6 100644 --- a/pgdog/src/backend/pool/pool_impl.rs +++ b/pgdog/src/backend/pool/pool_impl.rs @@ -6,18 +6,18 @@ use std::time::Duration; use futures::future::try_join_all; use once_cell::sync::{Lazy, OnceCell}; -use parking_lot::RwLock; -use parking_lot::{Mutex, RawMutex, lock_api::MutexGuard}; +use parking_lot::{Mutex, RawMutex, RwLock, lock_api::MutexGuard}; use pgdog_config::Role; use tokio::sync::Notify; use tokio::time::Instant; use tracing::{debug, error}; -use crate::backend::pool::LsnStats; -use crate::backend::{ConnectReason, DisconnectReason, Server, ServerOptions}; +use crate::backend::{ConnectReason, DisconnectReason, Server, ServerOptions, pool::LsnStats}; use crate::config::PoolerMode; -use crate::net::messages::{BackendPid, FrontendPid}; -use crate::net::{Liveness, Parameter, Parameters}; +use crate::net::{ + Liveness, Parameter, Parameters, + messages::{BackendPid, FrontendPid}, +}; use super::inner::CheckInResult; use super::{ @@ -397,7 +397,20 @@ impl Pool { /// Create a connection to the pool, untracked by the logic here. pub(crate) async fn standalone(&self, reason: ConnectReason) -> Result { - Monitor::create_connection(self, reason).await + let args = super::ConnectionArgs::from_pool(self, reason); + super::connection_creation::create(args).await + } + + /// Validate client token by attempting a connection to Postgres with the given token. + pub(crate) async fn validate_token(&self, token: &str) -> Result<(), Error> { + let addr = self.addr().clone().into_client_token(token); + let args = + super::ConnectionArgs::from_pool(self, ConnectReason::ValidateToken).with_addr(&addr); + super::connection_creation::create(args) + .await? + .disconnect_reason(DisconnectReason::CredentialsCheck); + + Ok(()) } /// Mark this pool offline and evict idle connections. diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 0b75d07a4..64d9af9a5 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -320,7 +320,11 @@ impl Server { // Perform authentication. let mut scram = Client::new(user, auth_secret); - let mut auth_type = AuthType::Trust; + let mut auth_type = if addr.server_auth.is_external_identity() { + AuthType::ExternalToken + } else { + AuthType::Trust + }; loop { let message = stream.read().await?; @@ -335,6 +339,7 @@ impl Server { match auth { Authentication::Ok => break, Authentication::ClearTextPassword => { + auth_type = AuthType::Plain; let password = Password::new_password(auth_secret.deref()); stream.send_flush(&password).await?; } diff --git a/pgdog/src/frontend/client/mod.rs b/pgdog/src/frontend/client/mod.rs index 8fc57ab9d..a1ce3ae79 100644 --- a/pgdog/src/frontend/client/mod.rs +++ b/pgdog/src/frontend/client/mod.rs @@ -16,6 +16,7 @@ use tracing::{Level as LogLevel, debug, enabled, error, info, trace, warn}; use super::{ClientRequest, Error, PreparedStatements}; use crate::auth::AuthResult; +use crate::auth::token_cache::AUTH_TOKEN_CACHE; use crate::auth::{md5, scram::Server}; use crate::backend::maintenance_mode; use crate::backend::pool::stats::MemoryStats; @@ -182,10 +183,11 @@ impl Client { async fn check_password( stream: &mut Stream, user: &str, + database: &str, auth_type: &AuthType, passwords: &[PasswordKind], ) -> Result { - if passwords.is_empty() { + if passwords.is_empty() && !auth_type.external_token() { return Ok(AuthResult::NoPasswordConfig); } @@ -241,6 +243,22 @@ impl Client { } AuthType::Trust => AuthResult::Ok, + + AuthType::ExternalToken => { + stream + .send_flush(&Authentication::ClearTextPassword) + .await?; + let response = stream.read().await?; + if let Some(password) = Password::from_bytes(response.to_bytes())?.password() { + if AUTH_TOKEN_CACHE.check(user, database, password).await? { + AuthResult::Ok + } else { + AuthResult::NoPasswordMatch + } + } else { + AuthResult::NoPasswordMatch + } + } }; Ok(result) @@ -282,7 +300,7 @@ impl Client { // The admin database is virtual and never present in the cluster // map, so authenticate directly against the configured admin password. let passwords = [PasswordKind::Plain(admin_password.clone())]; - Self::check_password(&mut stream, user, auth_type, &passwords).await? + Self::check_password(&mut stream, user, database, auth_type, &passwords).await? } else if passthrough { // Get the password. We always need it because we need to check if // it's current and hasn't been changed. @@ -328,7 +346,8 @@ impl Client { // entries to plaintext before the auth exchange let passwords = crate::auth::vault::resolve_passwords(cluster.passwords()).await; - Self::check_password(&mut stream, user, auth_type, &passwords).await? + Self::check_password(&mut stream, user, database, auth_type, &passwords) + .await? } }