use std::fmt::Debug;
use std::sync::Arc;
use async_trait::async_trait;
use futures::sink::Sink;
use pgwire::api::ClientInfo;
use pgwire::api::auth::StartupHandler;
use pgwire::error::{PgWireError, PgWireResult};
use pgwire::messages::{PgWireBackendMessage, PgWireFrontendMessage};
use crate::control::security::audit::{ArcAuditEmitter, AuditEvent};
use crate::control::security::credential::store::{AuthRejection, ScramLookup};
use crate::control::state::SharedState;
use super::super::handler::NodeDbPgHandler;
use super::provider::NodeDbParameterProvider;
pub(super) enum AuthStartup {
Trust(Arc<NodeDbPgHandler>),
Scram {
sasl: Box<pgwire::api::auth::sasl::SASLAuthStartupHandler<NodeDbParameterProvider>>,
state: Arc<SharedState>,
handler: Arc<NodeDbPgHandler>,
},
}
fn bind_startup_database<C: pgwire::api::ClientInfo>(
client: &C,
addr: &std::net::SocketAddr,
handler: &NodeDbPgHandler,
) {
let db_name = match client.metadata().get("database") {
Some(n) if !n.is_empty() => n.clone(),
_ => return,
};
handler.sessions.ensure_session(*addr);
let db_id = handler
.state
.credentials
.catalog()
.get_database_id_by_name(&db_name)
.ok()
.flatten();
if let Some(id) = db_id {
handler.sessions.set_current_database(addr, id);
}
}
#[async_trait]
impl StartupHandler for AuthStartup {
async fn on_startup<C>(
&self,
client: &mut C,
message: PgWireFrontendMessage,
) -> PgWireResult<()>
where
C: ClientInfo + futures::sink::Sink<PgWireBackendMessage> + Unpin + Send + Sync,
C::Error: Debug,
PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
{
match self {
AuthStartup::Trust(handler) => {
if let PgWireFrontendMessage::Startup(ref startup) = message {
pgwire::api::auth::protocol_negotiation(client, startup).await?;
pgwire::api::auth::save_startup_parameters_to_metadata(client, startup);
let identity = handler.resolve_trust_user(client)?;
let addr = client.socket_addr();
handler.sessions.ensure_session(addr);
handler.sessions.set_identity(&addr, identity);
pgwire::api::auth::finish_authentication(
client,
&super::provider::nodedb_parameter_provider(),
)
.await?;
}
let username = client
.metadata()
.get("user")
.cloned()
.unwrap_or_else(|| "unknown".to_string());
let source = client.socket_addr().to_string();
handler.state.audit_record(
AuditEvent::AuthSuccess,
None,
&source,
&format!("trust auth: {username}"),
);
let addr = client.socket_addr();
bind_startup_database(client, &addr, handler);
Ok(())
}
AuthStartup::Scram {
sasl,
state,
handler,
} => {
let was_in_auth = matches!(
client.state(),
pgwire::api::PgWireConnectionState::AuthenticationInProgress
);
let result = sasl.on_startup(client, message).await;
let username = client
.metadata()
.get("user")
.cloned()
.unwrap_or_else(|| "unknown".to_string());
let source = client.socket_addr().to_string();
match &result {
Ok(())
if was_in_auth
&& matches!(
client.state(),
pgwire::api::PgWireConnectionState::ReadyForQuery
) =>
{
state.credentials.record_login_success(&username);
state.audit_record(
AuditEvent::AuthSuccess,
None,
&source,
&format!("SCRAM-SHA-256 auth: {username}"),
);
let addr = client.socket_addr();
bind_startup_database(client, &addr, handler);
}
Err(_) if was_in_auth => {
let scram_ip_str = source
.parse::<std::net::SocketAddr>()
.map(|s| s.ip().to_string())
.unwrap_or_else(|_| source.clone());
let rate_limited = state
.rate_limiter
.is_login_rate_limited(&scram_ip_str, &username);
let counts = !rate_limited
&& matches!(
state.credentials.get_scram_credentials(&username),
ScramLookup::Found(_)
| ScramLookup::Rejected(AuthRejection::BadCredential)
);
if counts {
let emitter = ArcAuditEmitter(std::sync::Arc::clone(&state.audit));
let scram_ip =
source.parse::<std::net::SocketAddr>().ok().map(|s| s.ip());
state
.credentials
.record_login_failure(&username, scram_ip, &emitter);
state
.rate_limiter
.record_login_failure(&scram_ip_str, &username);
}
state.audit_record(
AuditEvent::AuthFailure,
None,
&source,
&format!("SCRAM-SHA-256 auth failed: {username}"),
);
}
_ => {}
}
result
}
}
}
}