diff --git a/src/bin/mirror.rs b/src/bin/mirror.rs index f3f9a3d..cb81baa 100644 --- a/src/bin/mirror.rs +++ b/src/bin/mirror.rs @@ -3,6 +3,7 @@ use clap::Parser; use reqwest::Url; use std::{net::SocketAddr, path::PathBuf}; use tokio::sync::mpsc; +use tokio::task::JoinSet; #[derive(Debug, clap::Args)] pub struct Args { @@ -53,36 +54,15 @@ pub async fn run( acme_directory_url, }: Args, ) -> anyhow::Result<()> { - let db = Db::new(wrap_pg.as_str(), wrap_pg_cert) - .await - .expect("to connect to pg for mirroring"); + let db = Db::new(wrap_pg.as_str(), wrap_pg_cert).await?; + + // TODO: allow starting up with polling backfill from beginning? + log::debug!("getting the latest op from the db..."); let latest = db .get_latest() - .await - .expect("to query for last createdAt") + .await? .expect("there to be at least one op in the db. did you backfill?"); - let (tx, rx) = mpsc::channel(2); - // upstream poller - let mut url = upstream.clone(); - tokio::task::spawn(async move { - log::info!("starting poll reader..."); - url.set_path("/export"); - tokio::task::spawn(async move { - poll_upstream(Some(latest), url, tx) - .await - .expect("to poll upstream for mirror sync") - }); - }); - // db writer - let poll_db = db.clone(); - tokio::task::spawn(async move { - log::info!("starting db writer..."); - pages_to_pg(poll_db, rx) - .await - .expect("to write to pg for mirror"); - }); - let listen_conf = match (bind, acme_domain.is_empty(), acme_cache_path) { (_, false, Some(cache_path)) => ListenConf::Acme { domains: acme_domain, @@ -93,9 +73,37 @@ pub async fn run( (_, _, _) => unreachable!(), }; - serve(&upstream, wrap, listen_conf) - .await - .expect("to be able to serve the mirror proxy app"); + let mut tasks = JoinSet::new(); + + let (send_page, recv_page) = mpsc::channel(8); + + let mut poll_url = upstream.clone(); + poll_url.set_path("/export"); + + tasks.spawn(poll_upstream(Some(latest), poll_url, send_page)); + tasks.spawn(pages_to_pg(db.clone(), recv_page)); + tasks.spawn(serve(upstream, wrap, listen_conf)); + + while let Some(next) = tasks.join_next().await { + match next { + Err(e) if e.is_panic() => { + log::error!("a joinset task panicked: {e}. bailing now. (should we panic?)"); + return Err(e.into()); + } + Err(e) => { + log::error!("a joinset task failed to join: {e}"); + return Err(e.into()); + } + Ok(Err(e)) => { + log::error!("a joinset task completed with error: {e}"); + return Err(e); + } + Ok(Ok(name)) => { + log::trace!("a task completed: {name:?}. {} left", tasks.len()); + } + } + } + Ok(()) } diff --git a/src/mirror.rs b/src/mirror.rs index 65c31fc..c4fdbe9 100644 --- a/src/mirror.rs +++ b/src/mirror.rs @@ -186,7 +186,9 @@ pub enum ListenConf { Bind(SocketAddr), } -pub async fn serve(upstream: &Url, plc: Url, listen: ListenConf) -> std::io::Result<()> { +pub async fn serve(upstream: Url, plc: Url, listen: ListenConf) -> anyhow::Result<&'static str> { + log::info!("starting server..."); + // not using crate CLIENT: don't want the retries etc let client = Client::builder() .user_agent(UA) @@ -231,11 +233,17 @@ pub async fn serve(upstream: &Url, plc: Url, listen: ListenConf) -> std::io::Res } let auto_cert = auto_cert.build().expect("acme config to build"); - run_insecure_notice(); - run(app, TcpListener::bind("0.0.0.0:443").acme(auto_cert)).await + let notice_task = tokio::task::spawn(run_insecure_notice()); + let app_res = run(app, TcpListener::bind("0.0.0.0:443").acme(auto_cert)).await; + log::warn!("server task ended, aborting insecure server task..."); + notice_task.abort(); + app_res?; + notice_task.await??; } - ListenConf::Bind(addr) => run(app, TcpListener::bind(addr)).await, + ListenConf::Bind(addr) => run(app, TcpListener::bind(addr)).await?, } + + Ok("server (uh oh?)") } async fn run(app: A, listener: L) -> std::io::Result<()> @@ -250,7 +258,7 @@ where } /// kick off a tiny little server on a tokio task to tell people to use 443 -fn run_insecure_notice() { +async fn run_insecure_notice() -> Result<(), std::io::Error> { #[handler] fn oop_plz_be_secure() -> (StatusCode, String) { ( @@ -266,11 +274,8 @@ You probably want to change your request to use HTTPS instead of HTTP. } let app = Route::new().at("/", get(oop_plz_be_secure)).with(Tracing); - let listener = TcpListener::bind("0.0.0.0:80"); - tokio::task::spawn(async move { - Server::new(listener) - .name("allegedly (mirror:80 helper)") - .run(app) - .await - }); + Server::new(TcpListener::bind("0.0.0.0:80")) + .name("allegedly (mirror:80 helper)") + .run(app) + .await } diff --git a/src/plc_pg.rs b/src/plc_pg.rs index ba861ea..ff86176 100644 --- a/src/plc_pg.rs +++ b/src/plc_pg.rs @@ -5,16 +5,15 @@ use std::path::PathBuf; use std::pin::pin; use std::time::Instant; use tokio::{ - task::{spawn, JoinHandle}, sync::{mpsc, oneshot}, + task::{JoinHandle, spawn}, }; use tokio_postgres::{ - Client, Error as PgError, NoTls, + Client, Error as PgError, NoTls, Socket, binary_copy::BinaryCopyInWriter, connect, - Socket, + tls::MakeTlsConnect, types::{Json, Type}, - tls::MakeTlsConnect }; fn get_tls(cert: PathBuf) -> anyhow::Result { @@ -86,7 +85,6 @@ impl Db { }) } - #[must_use] pub async fn connect(&self) -> Result<(Client, JoinHandle>), PgError> { log::trace!("connecting postgres..."); if let Some(ref connector) = self.cert { @@ -117,6 +115,8 @@ pub async fn pages_to_pg( db: Db, mut pages: mpsc::Receiver, ) -> anyhow::Result<&'static str> { + log::info!("starting pages_to_pg writer..."); + let (mut client, task) = db.connect().await?; let ops_stmt = client diff --git a/src/poll.rs b/src/poll.rs index efcae2f..c997af9 100644 --- a/src/poll.rs +++ b/src/poll.rs @@ -156,6 +156,7 @@ pub async fn poll_upstream( base: Url, dest: mpsc::Sender, ) -> anyhow::Result<&'static str> { + log::info!("starting upstream poller after {after:?}"); let mut tick = tokio::time::interval(UPSTREAM_REQUEST_INTERVAL); let mut prev_last: Option = after.map(Into::into); let mut boundary_state: Option = None;