diff --git a/Cargo.lock b/Cargo.lock index 105cf7d..27440d9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -38,7 +38,9 @@ dependencies = [ "governor", "http-body-util", "log", + "native-tls", "poem", + "postgres-native-tls", "reqwest", "reqwest-middleware", "reqwest-retry", @@ -1741,6 +1743,18 @@ version = "1.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f84267b20a16ea918e43c6a88433c2d54fa145c92a811b5b047ccbe153674483" +[[package]] +name = "postgres-native-tls" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1f39498473c92f7b6820ae970382c1d83178a3454c618161cb772e8598d9f6f" +dependencies = [ + "native-tls", + "tokio", + "tokio-native-tls", + "tokio-postgres", +] + [[package]] name = "postgres-protocol" version = "0.6.8" diff --git a/Cargo.toml b/Cargo.toml index 6d558c0..5cb91de 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,7 +15,9 @@ futures = "0.3.31" governor = "0.10.1" http-body-util = "0.1.3" log = "0.4.28" +native-tls = "0.2.14" poem = { version = "3.1.12", features = ["acme", "compression"] } +postgres-native-tls = "0.5.1" reqwest = { version = "0.12.23", features = ["stream", "json"] } reqwest-middleware = "0.4.2" reqwest-retry = "0.7.0" diff --git a/readme.md b/readme.md index b1e8406..96dcc6e 100644 --- a/readme.md +++ b/readme.md @@ -27,7 +27,7 @@ Allegedly can --upstream "https://plc.directory" \ --wrap "http://127.0.0.1:3000" \ --acme-domain "plc.wtf" \ - --acme-cache-dir ./acme-cache \ + --acme-cache-path ./acme-cache \ --acme-directory-url "https://acme-staging-v02.api.letsencrypt.org/directory" ``` diff --git a/src/bin/allegedly.rs b/src/bin/allegedly.rs index 0e230c4..3bd0f65 100644 --- a/src/bin/allegedly.rs +++ b/src/bin/allegedly.rs @@ -38,6 +38,8 @@ enum Commands { /// Pass a postgres connection url like "postgresql://localhost:5432" #[arg(long)] to_postgres: Option, + /// Cert for postgres (if needed) + postgres_cert: Option, /// Delete all operations from the postgres db before starting /// /// only used if `--to-postgres` is present @@ -79,6 +81,8 @@ enum Commands { /// the wrapped did-method-plc server's database (write access required) #[arg(long, env = "ALLEGEDLY_WRAP_PG")] wrap_pg: Url, + /// path to tls cert for the wrapped postgres db, if needed + wrap_pg_cert: Option, /// wrapping server listen address #[arg(short, long, env = "ALLEGEDLY_BIND")] #[clap(default_value = "127.0.0.1:8000")] @@ -166,6 +170,7 @@ async fn main() { dir, source_workers, to_postgres, + postgres_cert, postgres_reset, until, catch_up, @@ -195,9 +200,10 @@ async fn main() { }; let to_postgres_url_bulk = to_postgres.clone(); + let pg_cert = postgres_cert.clone(); let bulk_out_write = tokio::task::spawn(async move { if let Some(ref url) = to_postgres_url_bulk { - let db = Db::new(url.as_str()).await.unwrap(); + let db = Db::new(url.as_str(), pg_cert).await.unwrap(); backfill_to_pg(db, postgres_reset, rx, notify_last_at) .await .unwrap(); @@ -220,7 +226,7 @@ async fn main() { log::info!("writing catch-up pages"); let full_pages = full_pages(rx); if let Some(url) = to_postgres { - let db = Db::new(url.as_str()).await.unwrap(); + let db = Db::new(url.as_str(), postgres_cert).await.unwrap(); pages_to_pg(db, full_pages).await.unwrap(); } else { pages_to_stdout(full_pages, None).await.unwrap(); @@ -243,12 +249,13 @@ async fn main() { Commands::Mirror { wrap, wrap_pg, + wrap_pg_cert, bind, acme_domain, acme_cache_path, acme_directory_url, } => { - let db = Db::new(wrap_pg.as_str()).await.unwrap(); + let db = Db::new(wrap_pg.as_str(), wrap_pg_cert).await.unwrap(); let latest = db .get_latest() .await diff --git a/src/plc_pg.rs b/src/plc_pg.rs index 983d46b..55080b4 100644 --- a/src/plc_pg.rs +++ b/src/plc_pg.rs @@ -1,4 +1,7 @@ use crate::{Dt, ExportPage, Op, PageBoundaryState}; +use native_tls::{Certificate, TlsConnector}; +use postgres_native_tls::MakeTlsConnector; +use std::path::PathBuf; use std::pin::pin; use std::time::Instant; use tokio::sync::{mpsc, oneshot}; @@ -9,27 +12,54 @@ use tokio_postgres::{ types::{Json, Type}, }; +fn get_tls(cert: PathBuf) -> MakeTlsConnector { + let cert = std::fs::read(cert).unwrap(); + let cert = Certificate::from_pem(&cert).unwrap(); + let connector = TlsConnector::builder() + .add_root_certificate(cert) + .build() + .unwrap(); + MakeTlsConnector::new(connector) +} + /// a little tokio-postgres helper /// /// it's clone for easiness. it doesn't share any resources underneath after /// cloning at all so it's not meant for -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct Db { pg_uri: String, + cert: Option, } impl Db { - pub async fn new(pg_uri: &str) -> Result { + pub async fn new(pg_uri: &str, cert: Option) -> Result { // we're going to interact with did-method-plc's database, so make sure // it's what we expect: check for db migrations. log::trace!("checking migrations..."); - let (client, connection) = connect(pg_uri, NoTls).await?; - let connection_task = tokio::task::spawn(async move { - connection - .await - .inspect_err(|e| log::error!("connection ended with error: {e}")) - .unwrap(); - }); + + let connector = cert.map(get_tls); + + let (client, connection_task) = if let Some(ref connector) = connector { + let (client, connection) = connect(pg_uri, connector.clone()).await?; + let task = tokio::task::spawn(async move { + connection + .await + .inspect_err(|e| log::error!("connection ended with error: {e}")) + .unwrap(); + }); + (client, task) + } else { + let (client, connection) = connect(pg_uri, NoTls).await?; + let task = tokio::task::spawn(async move { + connection + .await + .inspect_err(|e| log::error!("connection ended with error: {e}")) + .unwrap(); + }); + (client, task) + }; + let migrations: Vec = client .query("SELECT name FROM kysely_migration ORDER BY name", &[]) .await? @@ -52,21 +82,37 @@ impl Db { Ok(Self { pg_uri: pg_uri.to_string(), + cert: connector, }) } pub async fn connect(&self) -> Result { log::trace!("connecting postgres..."); - let (client, connection) = connect(&self.pg_uri, NoTls).await?; + let client = if let Some(ref connector) = self.cert { + let (client, connection) = connect(&self.pg_uri, connector.clone()).await?; - // send the connection away to do the actual communication work - // apparently the connection will complete when the client drops - tokio::task::spawn(async move { - connection - .await - .inspect_err(|e| log::error!("connection ended with error: {e}")) - .unwrap(); - }); + // send the connection away to do the actual communication work + // apparently the connection will complete when the client drops + tokio::task::spawn(async move { + connection + .await + .inspect_err(|e| log::error!("connection ended with error: {e}")) + .unwrap(); + }); + client + } else { + let (client, connection) = connect(&self.pg_uri, NoTls).await?; + + // send the connection away to do the actual communication work + // apparently the connection will complete when the client drops + tokio::task::spawn(async move { + connection + .await + .inspect_err(|e| log::error!("connection ended with error: {e}")) + .unwrap(); + }); + client + }; Ok(client) }