diff options
Diffstat (limited to 'nym/src/main.rs')
| -rw-r--r-- | nym/src/main.rs | 320 |
1 files changed, 287 insertions, 33 deletions
diff --git a/nym/src/main.rs b/nym/src/main.rs index b906fc6..cc6e4a0 100644 --- a/nym/src/main.rs +++ b/nym/src/main.rs @@ -4,15 +4,26 @@ use nym_sdk::mixnet::{MixnetClientBuilder, MixnetMessageSender, Recipient, Stora use serde::{Deserialize, Serialize}; use std::env; use std::io::{self, Read}; -use std::path::PathBuf; +use std::os::unix::fs::{FileTypeExt, PermissionsExt}; +use std::path::{Path, PathBuf}; use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{UnixListener, UnixStream}; +use tokio::sync::mpsc; const MAX_PAYLOAD_BYTES: usize = 64 * 1024; +const MAX_REQUEST_BYTES: usize = 70 * 1024; const CONNECT_TIMEOUT: Duration = Duration::from_secs(120); const SEND_TIMEOUT: Duration = Duration::from_secs(60); +const REQUEST_TIMEOUT: Duration = Duration::from_secs(5); +const SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(240); +const DEFAULT_QUEUE_CAPACITY: usize = 8; +const MAX_QUEUE_CAPACITY: usize = 64; +const DEFAULT_SOCKET_PATH: &str = "/run/yamnweb/nym-sender.sock"; // send_plain_message queues the message for the background mixnet task. Keep // the client alive long enough for that task to flush the message before the -// process disconnects. +// one-shot compatibility process disconnects. Daemon mode keeps one client +// connected continuously, so it does not need a per-message grace period. const FLUSH_GRACE: Duration = Duration::from_secs(180); #[derive(Debug, Deserialize)] @@ -28,6 +39,34 @@ struct Response<'a> { error: Option<&'a str>, } +fn validate_request(request: &Request) -> Result<()> { + if request.payload.is_empty() || request.payload.len() > MAX_PAYLOAD_BYTES { + bail!("invalid envelope size"); + } + if request.payload.as_bytes().contains(&0) { + bail!("envelope contains a NUL byte"); + } + if request.entry_address.is_empty() + || request.entry_address.len() > 320 + || request.entry_address.chars().any(char::is_whitespace) + || request.entry_address.as_bytes().contains(&0) + { + bail!("invalid entry address"); + } + Ok(()) +} + +fn ingress_payload(request: &Request) -> Result<String> { + validate_request(request)?; + let raw_payload = BASE64.encode(request.payload.as_bytes()); + let request_payload = serde_json::json!({ + "version": 1, + "entry_address": request.entry_address, + "payload": raw_payload, + }); + serde_json::to_string(&request_payload).context("encode ingress envelope") +} + fn respond(response: Response<'_>) -> Result<()> { serde_json::to_writer(io::stdout(), &response).context("encode response")?; println!(); @@ -35,8 +74,7 @@ fn respond(response: Response<'_>) -> Result<()> { } fn configured_recipient() -> Result<Recipient> { - let value = env::var("YAMN_NYM_RECIPIENT") - .context("YAMN_NYM_RECIPIENT is not configured")?; + let value = env::var("YAMN_NYM_RECIPIENT").context("YAMN_NYM_RECIPIENT is not configured")?; value .parse::<Recipient>() .map_err(|_| anyhow::anyhow!("invalid configured Nym recipient")) @@ -50,26 +88,7 @@ fn storage_paths() -> Result<StoragePaths> { } async fn send(request: Request) -> Result<()> { - if request.payload.is_empty() || request.payload.len() > MAX_PAYLOAD_BYTES { - bail!("invalid envelope size"); - } - if request.payload.as_bytes().contains(&0) - || request - .entry_address - .chars() - .any(|c| matches!(c, '\r' | '\n' | '\0')) - { - bail!("envelope contains a NUL byte"); - } - - let raw_payload = BASE64 - .encode(request.payload.as_bytes()); - let request_payload = serde_json::json!({ - "version": 1, - "entry_address": request.entry_address, - "payload": raw_payload, - }); - let request_payload = serde_json::to_string(&request_payload).context("encode ingress envelope")?; + let request_payload = ingress_payload(&request)?; let recipient = configured_recipient()?; let paths = storage_paths()?; @@ -78,7 +97,7 @@ async fn send(request: Request) -> Result<()> { .context("prepare Nym client")? .build() .context("build Nym client")?; - let mut client = tokio::time::timeout(CONNECT_TIMEOUT, disconnected.connect_to_mixnet()) + let client = tokio::time::timeout(CONNECT_TIMEOUT, disconnected.connect_to_mixnet()) .await .context("Nym connection timed out")? .context("connect to Nym mixnet")?; @@ -96,17 +115,214 @@ async fn send(request: Request) -> Result<()> { Ok(()) } +fn queue_capacity() -> Result<usize> { + let Some(value) = env::var_os("YAMN_NYM_QUEUE_CAPACITY") else { + return Ok(DEFAULT_QUEUE_CAPACITY); + }; + let value = value + .to_str() + .context("YAMN_NYM_QUEUE_CAPACITY is not valid UTF-8")? + .parse::<usize>() + .context("YAMN_NYM_QUEUE_CAPACITY must be an integer")?; + if !(1..=MAX_QUEUE_CAPACITY).contains(&value) { + bail!("YAMN_NYM_QUEUE_CAPACITY must be between 1 and {MAX_QUEUE_CAPACITY}"); + } + Ok(value) +} + +fn socket_path() -> Result<PathBuf> { + let path = env::var_os("YAMN_NYM_SOCKET") + .map(PathBuf::from) + .unwrap_or_else(|| PathBuf::from(DEFAULT_SOCKET_PATH)); + if path.parent() != Some(Path::new("/run/yamnweb")) || path.file_name().is_none() { + bail!("YAMN_NYM_SOCKET must be directly inside /run/yamnweb"); + } + Ok(path) +} + +fn remove_stale_socket(path: &Path) -> Result<()> { + match std::fs::symlink_metadata(path) { + Ok(metadata) if metadata.file_type().is_socket() => { + std::fs::remove_file(path).context("remove stale Nym sender socket")?; + } + Ok(_) => bail!("refusing to replace a non-socket Nym sender path"), + Err(error) if error.kind() == io::ErrorKind::NotFound => {} + Err(error) => return Err(error).context("inspect Nym sender socket"), + } + Ok(()) +} + +struct SocketFile(PathBuf); + +impl Drop for SocketFile { + fn drop(&mut self) { + let _ = std::fs::remove_file(&self.0); + } +} + +async fn write_socket_response( + stream: &mut UnixStream, + success: bool, + error: Option<&'static str>, +) -> Result<()> { + let response = + serde_json::to_vec(&Response { success, error }).context("encode queue response")?; + stream + .write_all(&response) + .await + .context("write queue response")?; + stream + .write_all(b"\n") + .await + .context("finish queue response")?; + Ok(()) +} + +async fn accept_request(mut stream: UnixStream, sender: mpsc::Sender<Request>) -> Result<()> { + let mut input = Vec::new(); + tokio::time::timeout( + REQUEST_TIMEOUT, + (&mut stream) + .take((MAX_REQUEST_BYTES + 1) as u64) + .read_to_end(&mut input), + ) + .await + .context("queue request timed out")? + .context("read queue request")?; + if input.len() > MAX_REQUEST_BYTES { + write_socket_response(&mut stream, false, Some("Invalid Nym request")).await?; + return Ok(()); + } + let request: Request = match serde_json::from_slice(&input) { + Ok(request) => request, + Err(_) => { + write_socket_response(&mut stream, false, Some("Invalid Nym request")).await?; + return Ok(()); + } + }; + if validate_request(&request).is_err() { + write_socket_response(&mut stream, false, Some("Invalid Nym request")).await?; + return Ok(()); + } + match sender.try_send(request) { + Ok(()) => write_socket_response(&mut stream, true, None).await?, + Err(mpsc::error::TrySendError::Full(_)) => { + write_socket_response(&mut stream, false, Some("Nym queue is full")).await? + } + Err(mpsc::error::TrySendError::Closed(_)) => { + write_socket_response(&mut stream, false, Some("Nym sender is unavailable")).await? + } + } + Ok(()) +} + +async fn run_daemon() -> Result<()> { + let capacity = queue_capacity()?; + let path = socket_path()?; + let recipient = configured_recipient()?; + let paths = storage_paths()?; + let disconnected = MixnetClientBuilder::new_with_default_storage(paths) + .await + .context("prepare Nym client storage")? + .build() + .context("build Nym client")?; + let client = tokio::time::timeout(CONNECT_TIMEOUT, disconnected.connect_to_mixnet()) + .await + .context("Nym connection timed out")? + .context("connect to Nym mixnet")?; + + remove_stale_socket(&path)?; + let listener = UnixListener::bind(&path).context("bind Nym sender socket")?; + let _socket_file = SocketFile(path.clone()); + std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600)) + .context("restrict Nym sender socket permissions")?; + + let (sender, mut receiver) = mpsc::channel::<Request>(capacity); + let mut worker = tokio::spawn(async move { + while let Some(request) = receiver.recv().await { + let request_payload = ingress_payload(&request)?; + tokio::time::timeout( + SEND_TIMEOUT, + client.send_plain_message(recipient, request_payload), + ) + .await + .context("Nym send timed out")? + .context("send envelope through Nym")?; + } + tokio::time::sleep(FLUSH_GRACE).await; + client.disconnect().await; + Ok::<(), anyhow::Error>(()) + }); + + let mut terminate = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) + .context("install shutdown signal handler")?; + let mut worker_running = true; + let result: Result<()> = loop { + tokio::select! { + accepted = listener.accept() => { + let (stream, _) = accepted.context("accept Nym queue connection")?; + let request_sender = sender.clone(); + tokio::spawn(async move { + let _ = accept_request(stream, request_sender).await; + }); + } + signal = tokio::signal::ctrl_c() => { + signal.context("wait for shutdown signal")?; + break Ok(()); + } + _ = terminate.recv() => { + break Ok(()); + } + worker_result = &mut worker => { + worker_running = false; + break match worker_result { + Ok(Ok(())) => Err(anyhow::anyhow!("Nym sender worker stopped unexpectedly")), + Ok(Err(error)) => Err(error).context("process Nym sender queue"), + Err(error) => Err(error).context("join Nym sender worker"), + }; + } + } + }; + + drop(listener); + drop(sender); + if worker_running { + match tokio::time::timeout(SHUTDOWN_TIMEOUT, &mut worker).await { + Err(_) => { + worker.abort(); + let _ = worker.await; + return Err(anyhow::anyhow!( + "timed out draining the in-memory Nym queue" + )); + } + Ok(Err(error)) => return Err(error).context("join Nym sender worker"), + Ok(Ok(Err(error))) => return Err(error).context("process Nym sender queue"), + Ok(Ok(Ok(()))) => {} + } + } + result +} + +async fn run_oneshot() -> Result<()> { + let mut input = String::new(); + io::stdin() + .read_to_string(&mut input) + .context("read request")?; + let request: Request = serde_json::from_str(&input).context("decode request")?; + send(request).await +} + #[tokio::main] async fn main() { - let result = async { - let mut input = String::new(); - io::stdin() - .read_to_string(&mut input) - .context("read request")?; - let request: Request = serde_json::from_str(&input).context("decode request")?; - send(request).await + if env::args().nth(1).as_deref() == Some("--daemon") { + if let Err(error) = run_daemon().await { + eprintln!("yamn-nym-submit: daemon failed: {error:#}"); + std::process::exit(1); + } + return; } - .await; + + let result = run_oneshot().await; match result { Ok(()) => { @@ -125,3 +341,41 @@ async fn main() { } } } + +#[cfg(test)] +mod tests { + use super::*; + + fn valid_request() -> Request { + Request { + entry_address: "yamn@example.org".to_string(), + payload: "opaque envelope".to_string(), + } + } + + #[test] + fn validates_request_boundaries() { + assert!(validate_request(&valid_request()).is_ok()); + + let mut request = valid_request(); + request.entry_address = "yamn@example.org\r\nBcc: leak@example.org".to_string(); + assert!(validate_request(&request).is_err()); + + let mut request = valid_request(); + request.payload.clear(); + assert!(validate_request(&request).is_err()); + } + + #[test] + fn ingress_payload_contains_only_transport_fields() { + let encoded = ingress_payload(&valid_request()).expect("encode request"); + let value: serde_json::Value = serde_json::from_str(&encoded).expect("valid JSON"); + assert_eq!(value["version"], 1); + assert_eq!(value["entry_address"], "yamn@example.org"); + assert_eq!( + value["payload"], + BASE64.encode("opaque envelope".as_bytes()) + ); + assert_eq!(value.as_object().expect("object").len(), 3); + } +} |
