use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicI32, Ordering};
use std::time::Duration;
use anyhow::{Context as _, Result};
use dashmap::DashMap;
use rand::TryRngCore;
use surrealdb_core::ctx::CancelHandle;
use surrealdb_core::kvs::Datastore;
use tokio::net::TcpListener;
use tokio::sync::Semaphore;
use tokio_rustls::TlsAcceptor;
use tokio_util::sync::CancellationToken;
mod conn;
mod encode;
mod error;
mod msg;
mod sasl;
mod typing;
const LOG: &str = "surrealdb::pg";
pub(super) type CancelRegistry = DashMap<(i32, i32), CancelHandle>;
static NEXT_PID: AtomicI32 = AtomicI32::new(1);
fn next_backend_key() -> (i32, i32) {
let pid = NEXT_PID.fetch_add(1, Ordering::Relaxed);
let mut bytes = [0u8; 4];
let secret = if rand::rngs::OsRng.try_fill_bytes(&mut bytes).is_ok() {
i32::from_ne_bytes(bytes)
} else {
rand::random()
};
(pid, secret)
}
const MAX_CONNECTIONS: usize = 1024;
pub(crate) async fn start(
addr: SocketAddr,
ds: Arc<Datastore>,
ready: Arc<AtomicBool>,
shutdown: CancellationToken,
crt: Option<PathBuf>,
key: Option<PathBuf>,
) -> Result<()> {
let acceptor = match (crt, key) {
(Some(crt), Some(key)) => {
let config = axum_server::tls_rustls::RustlsConfig::from_pem_file(&crt, &key)
.await
.with_context(|| "Failed to load the postgres TLS certificate and key")?;
Some(Arc::new(TlsAcceptor::from(config.get_inner())))
}
_ => {
warn!(
target: LOG,
"The postgres wire protocol is serving plaintext connections; \
configure --web-crt/--web-key to enable TLS"
);
None
}
};
let listener = TcpListener::bind(addr)
.await
.with_context(|| format!("Failed to bind the postgres listener on {addr}"))?;
info!(target: LOG, "Started postgres wire protocol server on {}", addr);
tokio::spawn(accept_loop(listener, ds, ready, shutdown, acceptor));
Ok(())
}
async fn accept_loop(
listener: TcpListener,
ds: Arc<Datastore>,
ready: Arc<AtomicBool>,
shutdown: CancellationToken,
acceptor: Option<Arc<TlsAcceptor>>,
) {
let limiter = Arc::new(Semaphore::new(MAX_CONNECTIONS));
let registry: Arc<CancelRegistry> = Arc::new(DashMap::new());
loop {
tokio::select! {
biased;
_ = shutdown.cancelled() => break,
accepted = listener.accept() => match accepted {
Ok((stream, peer)) => {
if let Err(err) = stream.set_nodelay(true) {
debug!(target: LOG, "failed to set TCP_NODELAY for {peer}: {err}");
}
let Ok(permit) = Arc::clone(&limiter).try_acquire_owned() else {
warn!(target: LOG, "rejecting postgres connection from {peer}: too many connections");
tokio::spawn(conn::reject_overloaded(stream));
continue;
};
let (pid, secret) = next_backend_key();
tokio::spawn(conn::handle(
stream,
peer,
Arc::clone(&ds),
Arc::clone(&ready),
shutdown.clone(),
permit,
Arc::clone(®istry),
pid,
secret,
acceptor.clone(),
));
}
Err(err) => {
warn!(target: LOG, "failed to accept postgres connection: {err}");
tokio::time::sleep(Duration::from_millis(20)).await;
}
}
}
}
info!(target: LOG, "Stopped postgres wire protocol server");
}