#![allow(dead_code)]
use std::process::{Child, Command, ExitStatus, Stdio};
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
use powdb_query::executor::Engine;
use powdb_server::handler::{handle_connection, new_tx_gate, AuthRateLimiter, ConnOpts, TxGate};
use powdb_server::metrics::Metrics;
use powdb_server::protocol::Message;
use tokio::io::AsyncReadExt;
use tokio::net::TcpStream;
pub fn encode_connect(db: &str) -> Vec<u8> {
let mut payload = Vec::new();
payload.extend_from_slice(&(db.len() as u32).to_le_bytes());
payload.extend_from_slice(db.as_bytes());
payload.extend_from_slice(&0u32.to_le_bytes()); let mut frame = Vec::new();
frame.push(0x01); frame.push(0); frame.extend_from_slice(&(payload.len() as u32).to_le_bytes());
frame.extend_from_slice(&payload);
frame
}
pub fn encode_connect_with_password(db: &str, password: &str) -> Vec<u8> {
Message::Connect {
db_name: db.to_string(),
password: Some(zeroize::Zeroizing::new(password.to_string())),
username: None,
}
.encode()
}
pub fn encode_connect_user(db: &str, password: &str, username: &str) -> Vec<u8> {
Message::Connect {
db_name: db.to_string(),
password: Some(zeroize::Zeroizing::new(password.to_string())),
username: Some(username.to_string()),
}
.encode()
}
pub fn encode_query(q: &str) -> Vec<u8> {
let mut payload = Vec::new();
payload.extend_from_slice(&(q.len() as u32).to_le_bytes());
payload.extend_from_slice(q.as_bytes());
let mut frame = Vec::new();
frame.push(0x03); frame.push(0);
frame.extend_from_slice(&(payload.len() as u32).to_le_bytes());
frame.extend_from_slice(&payload);
frame
}
pub fn encode_ping() -> Vec<u8> {
vec![0x11, 0, 0, 0, 0, 0]
}
pub async fn read_response<S: AsyncReadExt + Unpin>(stream: &mut S) -> Vec<u8> {
let mut header = [0u8; 6];
stream.read_exact(&mut header).await.unwrap();
let payload_len = u32::from_le_bytes(header[2..6].try_into().unwrap()) as usize;
let mut payload = vec![0u8; payload_len];
if payload_len > 0 {
stream.read_exact(&mut payload).await.unwrap();
}
let mut full = Vec::with_capacity(6 + payload_len);
full.extend_from_slice(&header);
full.extend_from_slice(&payload);
full
}
pub async fn read_response_message<S: AsyncReadExt + Unpin>(stream: &mut S) -> Message {
Message::decode(&read_response(stream).await).unwrap()
}
pub async fn read_message<S: AsyncReadExt + Unpin>(stream: &mut S) -> Option<Message> {
let mut header = [0u8; 6];
match stream.read_exact(&mut header).await {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return None,
Err(_) => return None,
}
let payload_len = u32::from_le_bytes(header[2..6].try_into().unwrap()) as usize;
let mut payload = vec![0u8; payload_len];
if payload_len > 0 && stream.read_exact(&mut payload).await.is_err() {
return None;
}
let mut full = Vec::with_capacity(6 + payload_len);
full.extend_from_slice(&header);
full.extend_from_slice(&payload);
Message::decode(&full).ok()
}
pub fn fresh_engine() -> (Arc<RwLock<Engine>>, tempfile::TempDir) {
let tmp = tempfile::tempdir().unwrap();
let engine = Arc::new(RwLock::new(Engine::new(tmp.path()).unwrap()));
(engine, tmp)
}
pub fn unique_temp_dir(name: &str) -> std::path::PathBuf {
std::env::temp_dir().join(format!(
"powdb_{name}_{}_{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
))
}
pub struct InprocServer {
pub tx_wait_timeout: Duration,
pub idle_timeout: Duration,
pub query_timeout: Duration,
pub expected_password: Option<String>,
pub users: Arc<powdb_auth::UserStore>,
pub metrics: Arc<Metrics>,
pub rate_limiter: Option<AuthRateLimiter>,
pub shutdown_rx: Option<tokio::sync::watch::Receiver<bool>>,
pub single_conn: bool,
}
impl Default for InprocServer {
fn default() -> Self {
Self {
tx_wait_timeout: Duration::from_secs(5),
idle_timeout: Duration::from_secs(300),
query_timeout: Duration::from_secs(30),
expected_password: None,
users: Arc::new(powdb_auth::UserStore::new()),
metrics: Arc::new(Metrics::new()),
rate_limiter: None,
shutdown_rx: None,
single_conn: false,
}
}
}
async fn serve_one(
stream: TcpStream,
peer: std::net::SocketAddr,
engine: Arc<RwLock<Engine>>,
tx_gate: TxGate,
cfg: &InprocServer,
) {
let (own_tx, own_rx) = tokio::sync::watch::channel(false);
let _keep_shutdown_alive = own_tx;
let mut rx = cfg.shutdown_rx.clone().unwrap_or(own_rx);
handle_connection(
stream,
ConnOpts {
tx_wait_timeout: cfg.tx_wait_timeout,
db_name: None,
engine,
tx_gate,
expected_password: cfg.expected_password.clone().map(zeroize::Zeroizing::new),
users: cfg.users.clone(),
shutdown_rx: &mut rx,
idle_timeout: cfg.idle_timeout,
query_timeout: cfg.query_timeout,
rate_limiter: cfg.rate_limiter.as_ref(),
peer_addr: Some(peer),
metrics: cfg.metrics.clone(),
},
)
.await;
}
impl InprocServer {
pub async fn start(
self,
engine: Arc<RwLock<Engine>>,
) -> (std::net::SocketAddr, tokio::task::JoinHandle<()>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let tx_gate = new_tx_gate();
let cfg = self;
let handle = tokio::spawn(async move {
if cfg.single_conn {
let (stream, peer) = listener.accept().await.unwrap();
serve_one(stream, peer, engine, tx_gate, &cfg).await;
} else {
let cfg = Arc::new(cfg);
loop {
let (stream, peer) = match listener.accept().await {
Ok(v) => v,
Err(_) => break,
};
let engine = engine.clone();
let tx_gate = tx_gate.clone();
let cfg = cfg.clone();
tokio::spawn(async move {
serve_one(stream, peer, engine, tx_gate, &cfg).await;
});
}
}
});
(addr, handle)
}
}
pub fn free_port() -> u16 {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
listener.local_addr().unwrap().port()
}
pub fn spawn_server_bin_env(
port: u16,
data_dir: &std::path::Path,
extra_args: &[&str],
extra_env: &[(&str, &str)],
) -> Child {
let mut cmd = Command::new(env!("CARGO_BIN_EXE_powdb-server"));
cmd.arg("--port")
.arg(port.to_string())
.arg("--bind")
.arg("127.0.0.1")
.arg("--data-dir")
.arg(data_dir)
.args(extra_args)
.stdout(Stdio::null())
.stderr(Stdio::null());
for (key, value) in extra_env {
cmd.env(key, value);
}
cmd.spawn().expect("spawn powdb-server binary")
}
pub fn spawn_server_bin(port: u16, data_dir: &std::path::Path, extra_args: &[&str]) -> Child {
spawn_server_bin_env(port, data_dir, extra_args, &[])
}
pub async fn wait_for_bind(port: u16, within: Duration) -> TcpStream {
let start = Instant::now();
loop {
if let Ok(stream) = TcpStream::connect(("127.0.0.1", port)).await {
return stream;
}
assert!(
start.elapsed() <= within,
"server did not accept connections on port {port} within {within:?}"
);
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[cfg(unix)]
pub async fn wait_for_socket(path: &std::path::Path, within: Duration) -> tokio::net::UnixStream {
let start = Instant::now();
loop {
if let Ok(stream) = tokio::net::UnixStream::connect(path).await {
return stream;
}
assert!(
start.elapsed() <= within,
"server did not accept connections on socket {path:?} within {within:?}"
);
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[cfg(unix)]
pub fn send_sigterm(child: &Child) {
let status = Command::new("kill")
.arg("-TERM")
.arg(child.id().to_string())
.status()
.expect("invoke kill");
assert!(status.success(), "kill -TERM failed");
}
#[cfg(unix)]
pub fn send_sigkill(child: &Child) {
let status = Command::new("kill")
.arg("-9")
.arg(child.id().to_string())
.status()
.expect("invoke kill");
assert!(status.success(), "kill -9 failed");
}
pub fn wait_with_timeout(child: &mut Child, timeout: Duration) -> ExitStatus {
let start = Instant::now();
loop {
if let Some(status) = child.try_wait().expect("try_wait") {
return status;
}
if start.elapsed() > timeout {
let _ = child.kill();
panic!("server did not exit within {timeout:?}");
}
std::thread::sleep(Duration::from_millis(50));
}
}