use super::Write;
use super::outbound::Outbound;
use super::{Error, sealing};
use crate::LogId;
use darkbio_crypto::xhpke;
use std::fmt;
use std::sync::{Mutex, Weak};
use tracing::{debug, warn};
pub struct Sender<W: Write> {
outbound: Weak<Outbound<W>>, sealer: Weak<Mutex<xhpke::Sender>>, log_id: LogId, }
impl<W: Write> Sender<W> {
pub(super) fn new(
outbound: Weak<Outbound<W>>,
sealer: Weak<Mutex<xhpke::Sender>>,
log_id: LogId,
) -> Self {
Self {
outbound,
sealer,
log_id,
}
}
pub(crate) fn log_id(&self) -> LogId {
self.log_id
}
pub fn send(&self, message: &[u8]) -> Result<(), Error> {
let Some(outbound) = self.outbound.upgrade() else {
debug!("wire send refused, transport released");
return Err(Error::Terminated);
};
let Some(context) = self.sealer.upgrade() else {
debug!("wire send refused, session {} ended", self.log_id);
return Err(Error::EncryptionFailed("session ended".into()));
};
let mut sealer = context.lock().expect("encryption lock not poisoned");
let packet = match sealing::seal(&mut sealer, message) {
Ok(packet) => packet,
Err(Error::PacketTooLarge(size)) => {
warn!("wire message of {} bytes exceeds the sending limit", size);
return outbound.refuse_oversized(&context, size);
}
Err(err) => panic!("message encryption failed: {err}"),
};
let mut writer = outbound.lock();
drop(sealer);
let result = writer.send(&context, &packet, self.log_id);
if let Err(Error::EncryptionFailed(_)) = &result {
debug!("wire send refused, session {} ended", self.log_id);
}
result
}
pub(crate) fn disconnect(&self) -> Result<(), Error> {
if let (Some(outbound), Some(context)) = (self.outbound.upgrade(), self.sealer.upgrade()) {
outbound.disconnect(&context)?;
}
Ok(())
}
}
impl<W: Write> Clone for Sender<W> {
fn clone(&self) -> Self {
Self::new(self.outbound.clone(), self.sealer.clone(), self.log_id)
}
}
impl<W: Write> fmt::Debug for Sender<W> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Sender")
.field("session", &self.log_id)
.field("valid", &(self.sealer.strong_count() > 0))
.finish()
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
use crate::testing;
use crate::transport::Closer;
use crate::transport::DEFAULT_WRITE_TIMEOUT;
use crate::transport::framing::FrameReader;
use crate::transport::mock::payload;
use crate::transport::outbound::Side;
use crate::transport::testing::Memory;
use std::io;
use std::panic::{self, AssertUnwindSafe};
use std::sync::{Arc, TryLockError, mpsc};
use std::thread;
use std::time::{Duration, Instant};
fn connect<W: Write>(
outbound: &Arc<Outbound<W>>,
sender: xhpke::Sender,
) -> (Arc<Mutex<xhpke::Sender>>, Sender<W>) {
let sealer = Arc::new(Mutex::new(sender));
let sender = outbound.bind(&sealer);
(sealer, sender)
}
fn wait_sealing(sealer: &Mutex<xhpke::Sender>) {
let deadline = Instant::now() + Duration::from_secs(5);
while !matches!(sealer.try_lock(), Err(TryLockError::WouldBlock)) {
assert!(
Instant::now() < deadline,
"sender did not acquire encryption context"
);
thread::yield_now();
}
}
fn contexts() -> (xhpke::Sender, xhpke::Receiver) {
let secret = xhpke::SecretKey::generate();
let (sender, encap) = secret.public_key().new_sender(b"test").unwrap();
let receiver = secret.new_receiver(&encap, b"test").unwrap();
(sender, receiver)
}
#[derive(Clone, Default)]
struct Collector(Arc<Mutex<Vec<u8>>>);
impl Write for Collector {
fn set_write_deadline(&mut self, deadline: Instant) -> io::Result<()> {
testing::remaining(deadline)?;
Ok(())
}
}
impl io::Write for Collector {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
struct Gate {
entered: mpsc::Sender<()>,
release: Option<mpsc::Receiver<()>>,
dropped: mpsc::Sender<()>,
panics: bool,
deadline: Option<Instant>,
}
impl Gate {
fn new() -> (
Self,
mpsc::Receiver<()>,
mpsc::Sender<()>,
mpsc::Receiver<()>,
) {
let (entered_tx, entered) = mpsc::channel();
let (release, release_rx) = mpsc::channel();
let (dropped_tx, dropped) = mpsc::channel();
let gate = Self {
entered: entered_tx,
release: Some(release_rx),
dropped: dropped_tx,
panics: false,
deadline: None,
};
(gate, entered, release, dropped)
}
}
impl Write for Gate {
fn set_write_deadline(&mut self, deadline: Instant) -> io::Result<()> {
self.deadline = Some(deadline);
Ok(())
}
}
impl io::Write for Gate {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let deadline = self.deadline.expect("write deadline installed");
testing::remaining(deadline)?;
match self.release.take() {
Some(release) => {
let _ = self.entered.send(());
release
.recv_timeout(testing::remaining(deadline)?)
.map_err(|_| io::Error::from(io::ErrorKind::TimedOut))?;
if self.panics {
panic!("injected panic");
}
Err(io::Error::other("gate closed"))
}
None => Ok(buf.len()),
}
}
fn flush(&mut self) -> io::Result<()> {
testing::remaining(self.deadline.expect("write deadline installed"))?;
Ok(())
}
}
impl Drop for Gate {
fn drop(&mut self) {
let _ = self.dropped.send(());
}
}
#[test]
fn test_send_order() {
testing::init_tracing();
let (sender, mut receiver) = contexts();
let collector = Collector::default();
let outbound = Arc::new(Outbound::new(
collector.clone(),
Side::Client,
Closer::new(|| {}),
DEFAULT_WRITE_TIMEOUT,
));
let (_sealer, sender) = connect(&outbound, sender);
let threads: Vec<_> = (0..8)
.map(|thread| {
let sender = sender.clone();
thread::spawn(move || {
for i in 0..20 {
sender.send(&payload(thread * 100 + i)).unwrap();
}
})
})
.collect();
for thread in threads {
thread.join().unwrap();
}
let written = collector.0.lock().unwrap().clone();
let mut reader = FrameReader::new(Memory::new(&written[..]), Closer::new(|| {}));
let mut messages = Vec::new();
loop {
let packet = match reader.next_packet(None) {
Err(Error::Terminated) => break,
result => result.unwrap().unwrap(),
};
messages.push(sealing::open(&mut receiver, packet).unwrap());
}
messages.sort_unstable();
let mut expected: Vec<Vec<u8>> = (0..8)
.flat_map(|thread| (0..20).map(move |i| payload(thread * 100 + i)))
.collect();
expected.sort_unstable();
assert_eq!(messages, expected);
}
#[test]
fn test_end_with_queued_send() {
testing::init_tracing();
let collector = Collector::default();
let outbound = Arc::new(Outbound::new(
collector.clone(),
Side::Client,
Closer::new(|| {}),
DEFAULT_WRITE_TIMEOUT,
));
let (crypto, _) = contexts();
let (sealer, sender) = connect(&outbound, crypto);
let mut writer = outbound.lock();
let sending = thread::spawn(move || sender.send(&payload(1)));
wait_sealing(&sealer);
assert!(writer.end(&sealer));
drop(sealer);
drop(writer);
let (crypto, mut peer) = contexts();
let (_replacement, fresh) = connect(&outbound, crypto);
fresh.send(&payload(2)).unwrap();
assert!(matches!(
sending.join().unwrap(),
Err(Error::EncryptionFailed(_))
));
let bytes = collector.0.lock().unwrap().clone();
let mut reader = FrameReader::new(Memory::new(&bytes[..]), Closer::new(|| {}));
let packet = reader.next_packet(None).unwrap().unwrap();
assert_eq!(sealing::open(&mut peer, packet).unwrap(), payload(2));
assert!(matches!(reader.next_packet(None), Err(Error::Terminated)));
}
#[test]
fn test_send_failure_attribution() {
testing::init_tracing();
let (gate, entered, release, _) = Gate::new();
let (sender, _) = contexts();
let outbound = Arc::new(Outbound::new(
gate,
Side::Client,
Closer::new(|| {}),
DEFAULT_WRITE_TIMEOUT,
));
let (sealer, sender) = connect(&outbound, sender);
let first = {
let sender = sender.clone();
thread::spawn(move || sender.send(&payload(1)))
};
entered.recv_timeout(Duration::from_secs(5)).unwrap();
drop(sealer.try_lock().expect("encryption held during writing"));
let second = {
let sender = sender.clone();
thread::spawn(move || sender.send(&payload(2)))
};
wait_sealing(&sealer);
release.send(()).unwrap();
let result = first.join().unwrap();
assert!(matches!(result, Err(Error::SendFailed(_))), "{result:?}");
let result = second.join().unwrap();
assert!(
matches!(result, Err(Error::EncryptionFailed(_))),
"{result:?}"
);
let result = sender.send(&payload(3));
assert!(
matches!(result, Err(Error::EncryptionFailed(_))),
"{result:?}"
);
assert!(outbound.finish_receive(&sealer, Ok(Vec::new())).is_err());
}
#[test]
fn test_send_refusals() {
testing::init_tracing();
let outbound = Arc::new(Outbound::new(
Memory::new(Vec::new()),
Side::Client,
Closer::new(|| {}),
DEFAULT_WRITE_TIMEOUT,
));
let (crypto, _) = contexts();
let (first_sealer, first) = connect(&outbound, crypto);
first.send(&payload(1)).unwrap();
outbound.end(&first_sealer);
assert!(matches!(
first.send(&payload(2)),
Err(Error::EncryptionFailed(_))
));
let (crypto, _) = contexts();
let (second_sealer, second) = connect(&outbound, crypto);
assert!(matches!(
first.send(&payload(3)),
Err(Error::EncryptionFailed(_))
));
drop(first_sealer);
assert!(matches!(
first.send(&payload(4)),
Err(Error::EncryptionFailed(_))
));
second.send(&payload(5)).unwrap();
outbound.close();
outbound
.finish_receive(&second_sealer, Ok(Vec::new()))
.unwrap();
let result = second.send(&payload(6));
assert!(
matches!(&result, Err(Error::SendFailed(err)) if err.kind() == io::ErrorKind::NotConnected),
"{result:?}"
);
assert!(
outbound
.finish_receive(&second_sealer, Ok(Vec::new()))
.is_err()
);
drop(outbound);
assert!(matches!(second.send(&payload(7)), Err(Error::Terminated)));
}
#[test]
fn test_close_with_stuck_sends() {
testing::init_tracing();
let (gate, entered, release, dropped) = Gate::new();
let (sender, _) = contexts();
let closer = Closer::new(move || {
let _ = release.send(());
});
let outbound = Arc::new(Outbound::new(
gate,
Side::Client,
closer,
DEFAULT_WRITE_TIMEOUT,
));
let (sealer, sender) = connect(&outbound, sender);
let first = {
let sender = sender.clone();
thread::spawn(move || sender.send(&payload(1)))
};
entered.recv_timeout(Duration::from_secs(5)).unwrap();
let second = {
let sender = sender.clone();
thread::spawn(move || sender.send(&payload(2)))
};
wait_sealing(&sealer);
let (ending_tx, started) = mpsc::channel();
let ending = {
let outbound = outbound.clone();
let sealer = sealer.clone();
thread::spawn(move || {
ending_tx.send(()).unwrap();
outbound.end(&sealer);
})
};
started.recv_timeout(Duration::from_secs(5)).unwrap();
let (closed_tx, closed) = mpsc::channel();
let owner = {
let outbound = outbound.clone();
thread::spawn(move || {
outbound.close();
closed_tx.send(()).unwrap();
})
};
closed.recv_timeout(Duration::from_secs(5)).unwrap();
owner.join().unwrap();
ending.join().unwrap();
assert!(matches!(first.join().unwrap(), Err(Error::SendFailed(_))));
assert!(matches!(
second.join().unwrap(),
Err(Error::EncryptionFailed(_))
));
assert!(outbound.finish_receive(&sealer, Ok(Vec::new())).is_err());
assert!(matches!(
sender.send(&payload(3)),
Err(Error::EncryptionFailed(_))
));
drop(outbound);
dropped.recv_timeout(Duration::from_secs(5)).unwrap();
assert!(matches!(sender.send(&payload(4)), Err(Error::Terminated)));
}
#[test]
fn test_close_with_panicking_send() {
testing::init_tracing();
let (mut gate, entered, release, dropped) = Gate::new();
gate.panics = true;
let (sender, _) = contexts();
let closer = Closer::new(move || {
let _ = release.send(());
});
let outbound = Arc::new(Outbound::new(
gate,
Side::Client,
closer,
DEFAULT_WRITE_TIMEOUT,
));
let (_sealer, sender) = connect(&outbound, sender);
let sending = thread::spawn(move || {
panic::catch_unwind(AssertUnwindSafe(|| sender.send(&payload(1))))
});
entered.recv_timeout(Duration::from_secs(5)).unwrap();
outbound.close();
assert!(sending.join().unwrap().is_err());
drop(outbound);
dropped.recv_timeout(Duration::from_secs(5)).unwrap();
}
}