use std::path::PathBuf;
use std::process::ExitCode;
use std::io::{Read as _, Write as _};
use tokio::io::AsyncReadExt;
use weida::{
Address, Error, Identity, IncomingRequest, Publisher, Puller, Replier, Runtime, RuntimeConfig,
TransferMeta, Trust,
};
const CHUNK: usize = 64 * 1024;
#[derive(Clone, Copy, PartialEq, Eq)]
enum Framing {
Line,
Raw,
Nul,
Length,
}
impl Framing {
fn parse(name: &str) -> Framing {
match name {
"line" => Framing::Line,
"raw" => Framing::Raw,
"nul" => Framing::Nul,
"length" => Framing::Length,
other => fail(format!(
"--framing: {other:?} is not one of line, raw, nul, length"
)),
}
}
fn prefix(self, len: u64) -> Option<[u8; 8]> {
match self {
Framing::Length => Some(len.to_be_bytes()),
_ => None,
}
}
fn suffix(self) -> Option<u8> {
match self {
Framing::Line => Some(b'\n'),
Framing::Nul => Some(0),
Framing::Raw | Framing::Length => None,
}
}
}
const USAGE: &str = "usage:
weida send [--ca PATH] URL [FILE] one-way transfer; payload from FILE or stdin
weida request [--ca PATH] URL [FILE] exchange; the reply goes to stdout
weida sub [--ca PATH] [--filter F] [--count N] [--framing FORM] URL
weida serve (--echo | --sink | --pub) URL [--identity PATH] [--cert-out PATH]
options:
--ca PATH trust this certificate as an anchor instead of the address's key
--identity PATH load or create the server identity here, so the address is stable
--cert-out PATH write the server certificate, for a client that uses --ca
--filter F subscribe to this topic filter; the default takes every topic
--count N exit after N messages; without it, sub runs until stopped
--framing FORM how sub separates payloads on stdout: line (default), raw, nul, length
--version print the version
--help print this
URL is weida://[sha256:HEX@]HOST:PORT/PATH, weida+unix://SOCKET/PATH with the
socket path percent-encoded (weida+unix://%2Ftmp%2Fs.sock/echo), or
weida+pipe://NAME/PATH. weida+inproc:// is in-process only and cannot be
served from a shell.
Payload on stdin and stdout; addresses and diagnostics on stderr.
Exit codes: 2 usage, 3 refused, 4 unknown endpoint, 5 no reply, 6 untrusted,
7 connection lost, 1 anything else.";
fn usage() -> ! {
eprintln!("{USAGE}");
std::process::exit(2);
}
fn help() -> ! {
println!("{USAGE}");
std::process::exit(0);
}
fn fail(message: impl std::fmt::Display) -> ! {
eprintln!("weida: {message}");
std::process::exit(2);
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Role {
Echo,
Sink,
Publish,
}
enum Command {
Send {
url: String,
file: Option<PathBuf>,
ca: Option<PathBuf>,
},
Request {
url: String,
file: Option<PathBuf>,
ca: Option<PathBuf>,
},
Sub {
url: String,
filter: String,
count: Option<u64>,
framing: Framing,
ca: Option<PathBuf>,
},
Serve {
url: String,
role: Role,
identity: Option<PathBuf>,
cert_out: Option<PathBuf>,
},
}
fn parse_args() -> Command {
let mut args = std::env::args().skip(1);
let verb = args.next().unwrap_or_else(|| usage());
match verb.as_str() {
"--help" | "-h" | "help" => help(),
"--version" | "-V" => {
println!("weida {}", env!("CARGO_PKG_VERSION"));
std::process::exit(0);
}
_ => {}
}
let mut url = None;
let mut file = None;
let mut ca = None;
let mut identity = None;
let mut cert_out = None;
let mut filter = None;
let mut count = None;
let mut framing = Framing::Line;
let mut role = None;
let next = |args: &mut dyn Iterator<Item = String>, what: &str| -> String {
args.next()
.unwrap_or_else(|| fail(format!("{what} needs a value")))
};
while let Some(arg) = args.next() {
match arg.as_str() {
"--ca" => ca = Some(PathBuf::from(next(&mut args, "--ca"))),
"--identity" => identity = Some(PathBuf::from(next(&mut args, "--identity"))),
"--cert-out" => cert_out = Some(PathBuf::from(next(&mut args, "--cert-out"))),
"--filter" => filter = Some(next(&mut args, "--filter")),
"--count" => {
let value = next(&mut args, "--count");
count = Some(
value
.parse::<u64>()
.unwrap_or_else(|e| fail(format!("--count: {e}"))),
);
}
"--framing" => framing = Framing::parse(&next(&mut args, "--framing")),
"--echo" => role = Some(Role::Echo),
"--sink" => role = Some(Role::Sink),
"--pub" => role = Some(Role::Publish),
"--help" | "-h" => help(),
"--version" | "-V" => {
println!("weida {}", env!("CARGO_PKG_VERSION"));
std::process::exit(0);
}
other if other.starts_with("--") => fail(format!("unexpected option: {other}")),
other if url.is_none() => url = Some(other.to_owned()),
other if file.is_none() => file = Some(PathBuf::from(other)),
other => fail(format!("unexpected argument: {other}")),
}
}
let url = url.unwrap_or_else(|| usage());
match verb.as_str() {
"send" => Command::Send { url, file, ca },
"request" => Command::Request { url, file, ca },
"sub" => Command::Sub {
url,
filter: filter.unwrap_or_default(),
count,
framing,
ca,
},
"serve" => Command::Serve {
url,
role: role
.unwrap_or_else(|| fail("serve needs one of --echo, --sink or --pub".to_owned())),
identity,
cert_out,
},
other => fail(format!("unknown command: {other}")),
}
}
fn trust(ca: Option<&PathBuf>) -> Trust {
match ca {
Some(path) => Trust::anchor_file(path),
None => Trust::by_address(),
}
}
fn code_of(error: &Error) -> u8 {
match error {
Error::Rejected => 3,
Error::UnknownEndpoint => 4,
Error::NoReply => 5,
Error::Untrusted(_) => 6,
Error::ConnectionLost(_) | Error::NotConnected => 7,
_ => 1,
}
}
fn report(error: Error) -> ExitCode {
match &error {
Error::Untrusted(presented) => {
eprintln!("weida: the peer presented {presented}, which is not trusted");
}
other => eprintln!("weida: {other}"),
}
ExitCode::from(code_of(&error))
}
fn payload(file: Option<&PathBuf>) -> std::io::Result<Vec<u8>> {
match file {
Some(path) => std::fs::read(path),
None => {
let mut buf = Vec::new();
std::io::stdin().read_to_end(&mut buf)?;
Ok(buf)
}
}
}
#[tokio::main]
async fn main() -> ExitCode {
let command = parse_args();
let runtime = match Runtime::new(RuntimeConfig::default()) {
Ok(runtime) => runtime,
Err(e) => return report(e),
};
let outcome = match &command {
Command::Send { url, file, ca } => send(&runtime, url, file.as_ref(), ca.as_ref()).await,
Command::Request { url, file, ca } => {
request(&runtime, url, file.as_ref(), ca.as_ref()).await
}
Command::Sub {
url,
filter,
count,
framing,
ca,
} => subscribe(&runtime, url, filter, *count, *framing, ca.as_ref()).await,
Command::Serve {
url,
role,
identity,
cert_out,
} => serve(&runtime, url, *role, identity.as_ref(), cert_out.as_ref()).await,
};
let drained = runtime.drain(std::time::Duration::from_secs(5)).await;
if outcome.is_ok() && drained.outstanding > 0 {
eprintln!(
"weida: {} transfer(s) unacknowledged at the drain deadline",
drained.outstanding
);
}
match outcome {
Ok(()) => ExitCode::SUCCESS,
Err(e) => report(e),
}
}
async fn send(
runtime: &Runtime,
url: &str,
file: Option<&PathBuf>,
ca: Option<&PathBuf>,
) -> Result<(), Error> {
let body = payload(file).map_err(Error::Io)?;
let pusher = runtime.pusher(trust(ca));
pusher.connect(url).await?;
let mut transfer = pusher
.open(TransferMeta::default().with_content_len(body.len() as u64))
.await?;
transfer.write_all(&body).await?;
let delivery = transfer.finish()?;
match delivery.delivered().await {
Ok(()) => eprintln!("delivered bytes={}", body.len()),
Err(e) => {
eprintln!("undelivered bytes={}: {e}", body.len());
return Err(e);
}
}
Ok(())
}
async fn request(
runtime: &Runtime,
url: &str,
file: Option<&PathBuf>,
ca: Option<&PathBuf>,
) -> Result<(), Error> {
let body = payload(file).map_err(Error::Io)?;
let requester = runtime.requester(trust(ca));
requester.connect(url).await?;
let mut reply = requester
.request_with(
TransferMeta::default().with_content_len(body.len() as u64),
&body,
)
.await?;
let mut stdout = std::io::stdout();
let mut chunk = vec![0u8; CHUNK];
let mut total = 0u64;
loop {
let n = reply.read(&mut chunk).await?;
if n == 0 {
break;
}
total += n as u64;
stdout.write_all(&chunk[..n]).map_err(Error::Io)?;
}
stdout.flush().map_err(Error::Io)?;
eprintln!("reply bytes={total}");
Ok(())
}
async fn subscribe(
runtime: &Runtime,
url: &str,
filter: &str,
count: Option<u64>,
framing: Framing,
ca: Option<&PathBuf>,
) -> Result<(), Error> {
let subscriber = runtime.subscriber(trust(ca));
subscriber.connect(url).await?;
subscriber.subscribe(filter).await?;
eprintln!("subscribed filter={filter:?}");
let mut stdout = std::io::stdout();
let mut seen = 0u64;
let mut chunk = vec![0u8; CHUNK];
while count.is_none_or(|limit| seen < limit) {
let mut transfer = subscriber.recv().await?;
let topic = transfer.meta().topic.clone().unwrap_or_default();
let mut held = Vec::new();
let mut bytes = 0u64;
loop {
let n = transfer.read(&mut chunk).await?;
if n == 0 {
break;
}
bytes += n as u64;
if framing == Framing::Length {
held.extend_from_slice(&chunk[..n]);
} else {
stdout.write_all(&chunk[..n]).map_err(Error::Io)?;
}
}
if let Some(prefix) = framing.prefix(bytes) {
stdout.write_all(&prefix).map_err(Error::Io)?;
stdout.write_all(&held).map_err(Error::Io)?;
}
if let Some(suffix) = framing.suffix() {
stdout.write_all(&[suffix]).map_err(Error::Io)?;
}
stdout.flush().map_err(Error::Io)?;
match transfer.meta().gap.as_ref() {
Some(gap) => eprintln!("topic={topic} bytes={bytes} missed={}", gap.missed()),
None => eprintln!("topic={topic} bytes={bytes}"),
}
seen += 1;
}
Ok(())
}
async fn serve(
runtime: &Runtime,
url: &str,
role: Role,
identity: Option<&PathBuf>,
cert_out: Option<&PathBuf>,
) -> Result<(), Error> {
let listener = runtime.listener();
let (path, printed) = match url.strip_prefix("weida://") {
Some(rest) => {
let Some((authority, tail)) = rest.split_once('/') else {
return Err(Error::InvalidAddress(format!(
"missing endpoint path: {url:?}"
)));
};
if authority.contains('@') {
return Err(Error::InvalidAddress(
"a bind address carries no fingerprint; serve prints the one it used"
.to_owned(),
));
}
let socket: std::net::SocketAddr = authority
.parse()
.map_err(|e| Error::InvalidAddress(format!("{authority:?}: {e}")))?;
let identity = match identity {
Some(path) => load_or_create_identity(path)?,
None => Identity::generate_for(["localhost", "127.0.0.1", "::1"])?,
};
if let Some(path) = cert_out {
std::fs::write(path, identity.certificate_pem()?).map_err(Error::Io)?;
eprintln!("wrote the certificate to {}", path.display());
}
let fingerprint = identity.fingerprint()?;
let binding = listener.bind_quic(socket, identity).await?;
let local = binding.local_addr();
std::mem::forget(binding);
let path = format!("/{tail}");
let printed = format!(
"weida://{fingerprint}@{}:{}{path}",
local.ip(),
local.port()
);
(path, printed)
}
None => bind_local(&listener, url)?,
};
println!("{printed}");
match role {
Role::Echo => echo(listener.replier(&path)?).await,
Role::Sink => sink(listener.puller(&path)?).await,
Role::Publish => publish(listener.publisher(&path)?).await,
}
}
fn bind_local(listener: &weida::Listener, url: &str) -> Result<(String, String), Error> {
let address = Address::parse(url)?;
let path = address.path().to_owned();
match &address {
Address::Quic(_) => Err(Error::InvalidAddress(
"a weida:// bind address is handled above".to_owned(),
)),
Address::Inproc(addr) => Err(Error::InvalidAddress(format!(
"weida+inproc://{} is reachable only inside one process; \
use weida+unix://, weida+pipe:// or weida://",
addr.bus
))),
#[cfg(unix)]
Address::Unix(addr) => {
let binding = listener.bind_unix(&addr.socket)?;
std::mem::forget(binding);
Ok((path, addr.to_string()))
}
#[cfg(not(unix))]
Address::Unix(_) => Err(Error::InvalidAddress(
"weida+unix:// needs a Unix host".to_owned(),
)),
#[cfg(windows)]
Address::Pipe(addr) => {
let binding = listener.bind_pipe(&addr.name)?;
std::mem::forget(binding);
Ok((path, addr.to_string()))
}
#[cfg(not(windows))]
Address::Pipe(_) => Err(Error::InvalidAddress(
"weida+pipe:// needs a Windows host".to_owned(),
)),
}
}
async fn echo(replier: Replier) -> Result<(), Error> {
loop {
let request = replier.accept().await?;
tokio::spawn(async move {
if let Err(e) = echo_one(request).await {
eprintln!("weida: request failed: {e}");
}
});
}
}
async fn echo_one(mut request: IncomingRequest) -> Result<(), Error> {
let mut body = request.take_body();
let mut out = request.reply(TransferMeta::default()).await?;
let mut chunk = vec![0u8; CHUNK];
let mut total = 0u64;
loop {
let n = body.read(&mut chunk).await?;
if n == 0 {
break;
}
total += n as u64;
out.write_all(&chunk[..n]).await?;
}
out.finish()?;
eprintln!("echoed bytes={total}");
Ok(())
}
async fn sink(puller: Puller) -> Result<(), Error> {
let mut stdout = std::io::stdout();
let mut chunk = vec![0u8; CHUNK];
loop {
let mut transfer = puller.recv().await?;
let mut total = 0u64;
loop {
let n = transfer.read(&mut chunk).await?;
if n == 0 {
break;
}
total += n as u64;
stdout.write_all(&chunk[..n]).map_err(Error::Io)?;
}
stdout.flush().map_err(Error::Io)?;
eprintln!("received bytes={total}");
}
}
async fn publish(publisher: Publisher) -> Result<(), Error> {
let stdin = std::io::stdin();
let mut line = String::new();
loop {
line.clear();
let read = std::io::BufRead::read_line(&mut stdin.lock(), &mut line).map_err(Error::Io)?;
if read == 0 {
return Ok(());
}
let (topic, body) = match line.trim_end_matches('\n').split_once(' ') {
Some((topic, body)) => (topic, body),
None => ("", line.trim_end_matches('\n')),
};
let reached = publisher.publish(topic, body.as_bytes().to_vec())?;
eprintln!("published topic={topic} subscribers={reached}");
}
}
fn load_or_create_identity(path: &PathBuf) -> Result<Identity, Error> {
if path.exists() {
let identity = Identity::from_pem_file(path);
identity.fingerprint()?;
return Ok(identity);
}
let identity = Identity::generate_for(["localhost", "127.0.0.1", "::1"])?;
write_private(path, identity.to_pem()?.as_bytes()).map_err(Error::Io)?;
eprintln!("generated an identity in {}", path.display());
Ok(identity)
}
fn write_private(path: &PathBuf, bytes: &[u8]) -> std::io::Result<()> {
use std::io::Write as _;
let mut options = std::fs::OpenOptions::new();
options.write(true).create_new(true);
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt as _;
options.mode(0o600);
}
options.open(path)?.write_all(bytes)
}