1use super::Write;
8use super::outbound::Outbound;
9use super::{Error, sealing};
10use crate::LogId;
11use darkbio_crypto::xhpke;
12use std::fmt;
13use std::sync::{Mutex, Weak};
14use tracing::{debug, warn};
15
16pub struct Sender<W: Write> {
31 outbound: Weak<Outbound<W>>, sealer: Weak<Mutex<xhpke::Sender>>, log_id: LogId, }
35
36impl<W: Write> Sender<W> {
37 pub(super) fn new(
40 outbound: Weak<Outbound<W>>,
41 sealer: Weak<Mutex<xhpke::Sender>>,
42 log_id: LogId,
43 ) -> Self {
44 Self {
45 outbound,
46 sealer,
47 log_id,
48 }
49 }
50
51 pub(crate) fn log_id(&self) -> LogId {
53 self.log_id
54 }
55
56 pub fn send(&self, message: &[u8]) -> Result<(), Error> {
81 let Some(outbound) = self.outbound.upgrade() else {
82 debug!("wire send refused, transport released");
83 return Err(Error::Terminated);
84 };
85 let Some(context) = self.sealer.upgrade() else {
86 debug!("wire send refused, session {} ended", self.log_id);
87 return Err(Error::EncryptionFailed("session ended".into()));
88 };
89 let mut sealer = context.lock().expect("encryption lock not poisoned");
90 let packet = match sealing::seal(&mut sealer, message) {
91 Ok(packet) => packet,
92 Err(Error::PacketTooLarge(size)) => {
93 warn!("wire message of {} bytes exceeds the sending limit", size);
94 return outbound.refuse_oversized(&context, size);
95 }
96 Err(err) => panic!("message encryption failed: {err}"),
97 };
98 let mut writer = outbound.lock();
101 drop(sealer);
102 let result = writer.send(&context, &packet, self.log_id);
103 if let Err(Error::EncryptionFailed(_)) = &result {
104 debug!("wire send refused, session {} ended", self.log_id);
105 }
106 result
107 }
108
109 pub(crate) fn disconnect(&self) -> Result<(), Error> {
114 if let (Some(outbound), Some(context)) = (self.outbound.upgrade(), self.sealer.upgrade()) {
115 outbound.disconnect(&context)?;
116 }
117 Ok(())
118 }
119}
120
121impl<W: Write> Clone for Sender<W> {
122 fn clone(&self) -> Self {
123 Self::new(self.outbound.clone(), self.sealer.clone(), self.log_id)
124 }
125}
126
127impl<W: Write> fmt::Debug for Sender<W> {
128 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
130 f.debug_struct("Sender")
131 .field("session", &self.log_id)
132 .field("valid", &(self.sealer.strong_count() > 0))
133 .finish()
134 }
135}
136
137#[cfg(test)]
138#[cfg_attr(coverage_nightly, coverage(off))]
139mod tests {
140 use super::*;
141 use crate::testing;
142 use crate::transport::Closer;
143 use crate::transport::DEFAULT_WRITE_TIMEOUT;
144 use crate::transport::framing::FrameReader;
145 use crate::transport::mock::payload;
146 use crate::transport::outbound::Side;
147 use crate::transport::testing::Memory;
148 use std::io;
149 use std::panic::{self, AssertUnwindSafe};
150 use std::sync::{Arc, TryLockError, mpsc};
151 use std::thread;
152 use std::time::{Duration, Instant};
153
154 fn connect<W: Write>(
156 outbound: &Arc<Outbound<W>>,
157 sender: xhpke::Sender,
158 ) -> (Arc<Mutex<xhpke::Sender>>, Sender<W>) {
159 let sealer = Arc::new(Mutex::new(sender));
160 let sender = outbound.bind(&sealer);
161 (sealer, sender)
162 }
163
164 fn wait_sealing(sealer: &Mutex<xhpke::Sender>) {
167 let deadline = Instant::now() + Duration::from_secs(5);
168 while !matches!(sealer.try_lock(), Err(TryLockError::WouldBlock)) {
169 assert!(
170 Instant::now() < deadline,
171 "sender did not acquire encryption context"
172 );
173 thread::yield_now();
174 }
175 }
176
177 fn contexts() -> (xhpke::Sender, xhpke::Receiver) {
179 let secret = xhpke::SecretKey::generate();
180 let (sender, encap) = secret.public_key().new_sender(b"test").unwrap();
181 let receiver = secret.new_receiver(&encap, b"test").unwrap();
182 (sender, receiver)
183 }
184
185 #[derive(Clone, Default)]
187 struct Collector(Arc<Mutex<Vec<u8>>>);
188
189 impl Write for Collector {
190 fn set_write_deadline(&mut self, deadline: Instant) -> io::Result<()> {
191 testing::remaining(deadline)?;
192 Ok(())
193 }
194 }
195
196 impl io::Write for Collector {
197 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
198 self.0.lock().unwrap().extend_from_slice(buf);
199 Ok(buf.len())
200 }
201
202 fn flush(&mut self) -> io::Result<()> {
203 Ok(())
204 }
205 }
206
207 struct Gate {
210 entered: mpsc::Sender<()>,
211 release: Option<mpsc::Receiver<()>>,
212 dropped: mpsc::Sender<()>,
213 panics: bool,
214 deadline: Option<Instant>,
215 }
216
217 impl Gate {
218 fn new() -> (
221 Self,
222 mpsc::Receiver<()>,
223 mpsc::Sender<()>,
224 mpsc::Receiver<()>,
225 ) {
226 let (entered_tx, entered) = mpsc::channel();
227 let (release, release_rx) = mpsc::channel();
228 let (dropped_tx, dropped) = mpsc::channel();
229 let gate = Self {
230 entered: entered_tx,
231 release: Some(release_rx),
232 dropped: dropped_tx,
233 panics: false,
234 deadline: None,
235 };
236 (gate, entered, release, dropped)
237 }
238 }
239
240 impl Write for Gate {
241 fn set_write_deadline(&mut self, deadline: Instant) -> io::Result<()> {
242 self.deadline = Some(deadline);
243 Ok(())
244 }
245 }
246
247 impl io::Write for Gate {
248 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
249 let deadline = self.deadline.expect("write deadline installed");
250 testing::remaining(deadline)?;
251 match self.release.take() {
252 Some(release) => {
253 let _ = self.entered.send(());
254 release
255 .recv_timeout(testing::remaining(deadline)?)
256 .map_err(|_| io::Error::from(io::ErrorKind::TimedOut))?;
257 if self.panics {
258 panic!("injected panic");
259 }
260 Err(io::Error::other("gate closed"))
261 }
262 None => Ok(buf.len()),
263 }
264 }
265
266 fn flush(&mut self) -> io::Result<()> {
267 testing::remaining(self.deadline.expect("write deadline installed"))?;
268 Ok(())
269 }
270 }
271
272 impl Drop for Gate {
273 fn drop(&mut self) {
274 let _ = self.dropped.send(());
275 }
276 }
277
278 #[test]
281 fn test_send_order() {
282 testing::init_tracing();
283
284 let (sender, mut receiver) = contexts();
285 let collector = Collector::default();
286 let outbound = Arc::new(Outbound::new(
287 collector.clone(),
288 Side::Client,
289 Closer::new(|| {}),
290 DEFAULT_WRITE_TIMEOUT,
291 ));
292 let (_sealer, sender) = connect(&outbound, sender);
293
294 let threads: Vec<_> = (0..8)
295 .map(|thread| {
296 let sender = sender.clone();
297 thread::spawn(move || {
298 for i in 0..20 {
299 sender.send(&payload(thread * 100 + i)).unwrap();
300 }
301 })
302 })
303 .collect();
304 for thread in threads {
305 thread.join().unwrap();
306 }
307
308 let written = collector.0.lock().unwrap().clone();
310 let mut reader = FrameReader::new(Memory::new(&written[..]), Closer::new(|| {}));
311 let mut messages = Vec::new();
312 loop {
313 let packet = match reader.next_packet(None) {
314 Err(Error::Terminated) => break,
315 result => result.unwrap().unwrap(),
316 };
317 messages.push(sealing::open(&mut receiver, packet).unwrap());
318 }
319 messages.sort_unstable();
320 let mut expected: Vec<Vec<u8>> = (0..8)
321 .flat_map(|thread| (0..20).map(move |i| payload(thread * 100 + i)))
322 .collect();
323 expected.sort_unstable();
324 assert_eq!(messages, expected);
325 }
326
327 #[test]
331 fn test_end_with_queued_send() {
332 testing::init_tracing();
333
334 let collector = Collector::default();
335 let outbound = Arc::new(Outbound::new(
336 collector.clone(),
337 Side::Client,
338 Closer::new(|| {}),
339 DEFAULT_WRITE_TIMEOUT,
340 ));
341 let (crypto, _) = contexts();
342 let (sealer, sender) = connect(&outbound, crypto);
343 let mut writer = outbound.lock();
344 let sending = thread::spawn(move || sender.send(&payload(1)));
345
346 wait_sealing(&sealer);
347 assert!(writer.end(&sealer));
348 drop(sealer);
349 drop(writer);
350
351 let (crypto, mut peer) = contexts();
352 let (_replacement, fresh) = connect(&outbound, crypto);
353 fresh.send(&payload(2)).unwrap();
354 assert!(matches!(
355 sending.join().unwrap(),
356 Err(Error::EncryptionFailed(_))
357 ));
358
359 let bytes = collector.0.lock().unwrap().clone();
360 let mut reader = FrameReader::new(Memory::new(&bytes[..]), Closer::new(|| {}));
361 let packet = reader.next_packet(None).unwrap().unwrap();
362 assert_eq!(sealing::open(&mut peer, packet).unwrap(), payload(2));
363 assert!(matches!(reader.next_packet(None), Err(Error::Terminated)));
364 }
365
366 #[test]
370 fn test_send_failure_attribution() {
371 testing::init_tracing();
372
373 let (gate, entered, release, _) = Gate::new();
374 let (sender, _) = contexts();
375 let outbound = Arc::new(Outbound::new(
376 gate,
377 Side::Client,
378 Closer::new(|| {}),
379 DEFAULT_WRITE_TIMEOUT,
380 ));
381 let (sealer, sender) = connect(&outbound, sender);
382
383 let first = {
386 let sender = sender.clone();
387 thread::spawn(move || sender.send(&payload(1)))
388 };
389 entered.recv_timeout(Duration::from_secs(5)).unwrap();
390 drop(sealer.try_lock().expect("encryption held during writing"));
392 let second = {
393 let sender = sender.clone();
394 thread::spawn(move || sender.send(&payload(2)))
395 };
396 wait_sealing(&sealer);
397
398 release.send(()).unwrap();
400 let result = first.join().unwrap();
401 assert!(matches!(result, Err(Error::SendFailed(_))), "{result:?}");
402 let result = second.join().unwrap();
403 assert!(
404 matches!(result, Err(Error::EncryptionFailed(_))),
405 "{result:?}"
406 );
407 let result = sender.send(&payload(3));
408 assert!(
409 matches!(result, Err(Error::EncryptionFailed(_))),
410 "{result:?}"
411 );
412 assert!(outbound.finish_receive(&sealer, Ok(Vec::new())).is_err());
413 }
414
415 #[test]
419 fn test_send_refusals() {
420 testing::init_tracing();
421
422 let outbound = Arc::new(Outbound::new(
423 Memory::new(Vec::new()),
424 Side::Client,
425 Closer::new(|| {}),
426 DEFAULT_WRITE_TIMEOUT,
427 ));
428 let (crypto, _) = contexts();
429 let (first_sealer, first) = connect(&outbound, crypto);
430 first.send(&payload(1)).unwrap();
431
432 outbound.end(&first_sealer);
433 assert!(matches!(
434 first.send(&payload(2)),
435 Err(Error::EncryptionFailed(_))
436 ));
437
438 let (crypto, _) = contexts();
439 let (second_sealer, second) = connect(&outbound, crypto);
440 assert!(matches!(
441 first.send(&payload(3)),
442 Err(Error::EncryptionFailed(_))
443 ));
444 drop(first_sealer);
445 assert!(matches!(
446 first.send(&payload(4)),
447 Err(Error::EncryptionFailed(_))
448 ));
449 second.send(&payload(5)).unwrap();
450
451 outbound.close();
453 outbound
454 .finish_receive(&second_sealer, Ok(Vec::new()))
455 .unwrap();
456 let result = second.send(&payload(6));
457 assert!(
458 matches!(&result, Err(Error::SendFailed(err)) if err.kind() == io::ErrorKind::NotConnected),
459 "{result:?}"
460 );
461 assert!(
462 outbound
463 .finish_receive(&second_sealer, Ok(Vec::new()))
464 .is_err()
465 );
466 drop(outbound);
467 assert!(matches!(second.send(&payload(7)), Err(Error::Terminated)));
468 }
469
470 #[test]
475 fn test_close_with_stuck_sends() {
476 testing::init_tracing();
477
478 let (gate, entered, release, dropped) = Gate::new();
479 let (sender, _) = contexts();
480 let closer = Closer::new(move || {
481 let _ = release.send(());
482 });
483 let outbound = Arc::new(Outbound::new(
484 gate,
485 Side::Client,
486 closer,
487 DEFAULT_WRITE_TIMEOUT,
488 ));
489 let (sealer, sender) = connect(&outbound, sender);
490
491 let first = {
492 let sender = sender.clone();
493 thread::spawn(move || sender.send(&payload(1)))
494 };
495 entered.recv_timeout(Duration::from_secs(5)).unwrap();
496 let second = {
497 let sender = sender.clone();
498 thread::spawn(move || sender.send(&payload(2)))
499 };
500 wait_sealing(&sealer);
501
502 let (ending_tx, started) = mpsc::channel();
503 let ending = {
504 let outbound = outbound.clone();
505 let sealer = sealer.clone();
506 thread::spawn(move || {
507 ending_tx.send(()).unwrap();
508 outbound.end(&sealer);
509 })
510 };
511 started.recv_timeout(Duration::from_secs(5)).unwrap();
512
513 let (closed_tx, closed) = mpsc::channel();
514 let owner = {
515 let outbound = outbound.clone();
516 thread::spawn(move || {
517 outbound.close();
518 closed_tx.send(()).unwrap();
519 })
520 };
521 closed.recv_timeout(Duration::from_secs(5)).unwrap();
522 owner.join().unwrap();
523 ending.join().unwrap();
524 assert!(matches!(first.join().unwrap(), Err(Error::SendFailed(_))));
525 assert!(matches!(
526 second.join().unwrap(),
527 Err(Error::EncryptionFailed(_))
528 ));
529 assert!(outbound.finish_receive(&sealer, Ok(Vec::new())).is_err());
530 assert!(matches!(
531 sender.send(&payload(3)),
532 Err(Error::EncryptionFailed(_))
533 ));
534 drop(outbound);
535 dropped.recv_timeout(Duration::from_secs(5)).unwrap();
536 assert!(matches!(sender.send(&payload(4)), Err(Error::Terminated)));
537 }
538
539 #[test]
542 fn test_close_with_panicking_send() {
543 testing::init_tracing();
544
545 let (mut gate, entered, release, dropped) = Gate::new();
546 gate.panics = true;
547 let (sender, _) = contexts();
548 let closer = Closer::new(move || {
549 let _ = release.send(());
550 });
551 let outbound = Arc::new(Outbound::new(
552 gate,
553 Side::Client,
554 closer,
555 DEFAULT_WRITE_TIMEOUT,
556 ));
557 let (_sealer, sender) = connect(&outbound, sender);
558
559 let sending = thread::spawn(move || {
560 panic::catch_unwind(AssertUnwindSafe(|| sender.send(&payload(1))))
561 });
562 entered.recv_timeout(Duration::from_secs(5)).unwrap();
563 outbound.close();
564 assert!(sending.join().unwrap().is_err());
565 drop(outbound);
566 dropped.recv_timeout(Duration::from_secs(5)).unwrap();
567 }
568}