1use std::{net::SocketAddr, sync::Arc, task::Waker};
2
3use crate::{IdBytes, Result, cenc::validate_id};
4use fnv::FnvHashMap;
5use futures::{
6 Sink, Stream,
7 task::{Context, Poll},
8};
9use rand::Rng;
10use std::{
11 collections::VecDeque,
12 io,
13 pin::Pin,
14 sync::atomic::{AtomicU16, Ordering},
15 time::Duration,
16};
17use tokio::sync::oneshot::{self, Receiver, Sender};
18use tracing::{error, trace};
19use wasm_timer::Instant;
20
21use super::{
22 Command, Peer, QueryAndTid,
23 cenc::{generic_hash, generic_hash_with_key},
24 message::{MsgData, ReplyMsgData, RequestMsgData},
25 query::QueryId,
26 stateobserver::Observer,
27 stream::MessageDataStream,
28 thirty_two_random_bytes,
29};
30
31const ROTATE_INTERVAL: u64 = 300_000;
32
33pub type Tid = u16;
34
35#[derive(Debug)]
37pub struct Secrets {
38 _rotate_counter: usize,
39 secrets: [[u8; 32]; 2],
43 _rotation: Duration,
44 _last_rotation: Instant,
45}
46
47impl Default for Secrets {
48 fn default() -> Self {
49 Self {
50 _rotate_counter: 10,
51 _rotation: Duration::from_millis(ROTATE_INTERVAL),
52 _last_rotation: Instant::now(),
53 secrets: [thirty_two_random_bytes(), thirty_two_random_bytes()],
54 }
55 }
56}
57
58impl Secrets {
59 fn _rotate_secrets(&mut self) -> Result<()> {
60 let tmp = self.secrets[0];
61 self.secrets[0] = self.secrets[1];
62 self.secrets[1] = generic_hash(&tmp);
63 Ok(())
64 }
65
66 fn _drain(&mut self) -> Result<()> {
67 self._rotate_counter -= 1;
68 if self._rotate_counter == 0 {
69 self._rotate_counter = 10;
70 self._rotate_secrets()?;
71 }
72 Ok(())
73 }
74
75 pub fn token(&self, peer: &Peer, secret_index: usize) -> Result<[u8; 32]> {
76 generic_hash_with_key(
77 &peer.socketv4()?.ip().octets()[..],
78 &self.secrets[secret_index],
79 )
80 }
81}
82
83#[derive(Debug, Clone)]
85pub struct InResponse {
86 pub request: Box<RequestMsgData>,
87 pub response: ReplyMsgData,
88 pub peer: Peer,
90 pub query_id: Option<QueryId>,
91}
92
93impl InResponse {
94 pub fn tid(&self) -> Tid {
95 self.request.tid
96 }
97 pub fn cmd(&self) -> Command {
98 self.request.command
99 }
100
101 pub fn valid_peer_id(&self) -> Option<IdBytes> {
102 validate_id(&self.response.id, &self.peer)
103 }
104
105 fn new(
106 request: Box<RequestMsgData>,
107 response: ReplyMsgData,
108 peer: Peer,
109 query_id: Option<QueryId>,
110 ) -> Self {
111 Self {
112 request,
113 response,
114 peer,
115 query_id,
116 }
117 }
118}
119
120#[derive(Debug)]
121pub struct OutRequestBuilder {
122 peer: Peer,
123 command: Command,
124 tid: Option<u16>,
125 id: Option<[u8; 32]>,
126 query_id: Option<QueryId>,
127 token: Option<[u8; 32]>,
128 target: Option<IdBytes>,
129 value: Option<Vec<u8>>,
130}
131
132macro_rules! setter {
133 ($name:ident, $type:ty) => {
134 pub fn $name(mut self, $name: $type) -> Self {
135 self.$name = Some($name);
136 self
137 }
138 };
139}
140impl OutRequestBuilder {
141 pub fn from_request(req: RequestMsgData) -> Self {
142 Self {
143 peer: req.to.clone(),
144 command: req.command,
145 tid: Some(req.tid),
146 id: req.id,
147 query_id: None,
148 token: req.token,
149 target: req.target.map(IdBytes::from),
150 value: req.value,
151 }
152 }
153 pub fn new(peer: Peer, command: Command) -> Self {
154 Self {
155 peer,
156 command,
157 tid: None,
158 query_id: None,
159 token: None,
160 target: None,
161 value: None,
162 id: None,
163 }
164 }
165 pub fn peer(mut self, peer: Peer) -> Self {
166 self.peer = peer;
167 self
168 }
169 setter!(tid, u16);
170 setter!(query_id, QueryId);
171 setter!(token, [u8; 32]);
172 setter!(target, IdBytes);
173 setter!(value, Vec<u8>);
174}
175
176#[derive(Debug)]
178pub enum OutMessage {
179 Request((Option<QueryId>, RequestMsgData, Option<Sender<()>>)),
180 Reply((Option<Sender<()>>, ReplyMsgData)),
181}
182
183impl OutMessage {
184 fn into_sendable(self) -> (MsgData, SocketAddr, Option<Sender<()>>) {
185 match self {
186 OutMessage::Request((_query_id, msg, tx)) => {
187 let dest = SocketAddr::from(&msg.to);
188 (MsgData::Request(msg), dest, tx)
189 }
190 OutMessage::Reply((tx, msg)) => {
191 let dest = SocketAddr::from(&msg.to);
192 (MsgData::Reply(msg), dest, tx)
193 }
194 }
195 }
196}
197
198#[derive(Debug)]
199struct InflightRequest {
200 message: RequestMsgData,
202 #[expect(unused)] timestamp: Instant,
205 query_id: Option<QueryId>,
207}
208
209#[derive(Debug)]
210pub struct IoHandler {
211 id: Observer<IdBytes>,
212 ephemeral: bool,
213 message_stream: MessageDataStream,
214 pending_send: VecDeque<OutMessage>,
216 pending_flush: Option<OutMessage>,
218 pending_recv: FnvHashMap<Tid, InflightRequest>,
220 secrets: Secrets,
221 tid: AtomicU16,
222 stream_waker: Option<Waker>,
223 name: String,
224}
225
226impl IoHandler {
227 pub fn new(id: Observer<IdBytes>, message_stream: MessageDataStream, config: IoConfig) -> Self {
228 Self {
229 id,
230 ephemeral: config.ephemeral,
231 message_stream,
232 pending_send: Default::default(),
233 pending_flush: None,
234 pending_recv: Default::default(),
235 secrets: config.secrets,
236 tid: AtomicU16::new(rand::thread_rng().r#gen()),
237 stream_waker: Default::default(),
238 name: random_name(),
239 }
240 }
241 pub fn name(&self) -> &str {
242 &self.name
243 }
244
245 pub fn is_ephemeral(&self) -> bool {
246 self.ephemeral
247 }
248
249 pub fn id(&self) -> IdBytes {
250 *self.id.get()
251 }
252
253 pub fn local_addr(&self) -> crate::Result<SocketAddr> {
254 self.message_stream.local_addr()
255 }
256 pub fn socket(&self) -> udx::UdxSocket {
257 self.message_stream.socket()
258 }
259 pub fn token(&self, peer: &Peer, secret_index: usize) -> crate::Result<[u8; 32]> {
261 self.secrets.token(peer, secret_index)
262 }
263
264 pub fn new_tid(&self) -> Tid {
265 self.tid.fetch_add(1, Ordering::Relaxed)
266 }
267
268 pub fn enqueue_reply(&mut self, msg: ReplyMsgData, tx: Option<Sender<()>>) {
269 self.pending_send.push_back(OutMessage::Reply((tx, msg)));
270 self.maybe_wake();
271 }
272
273 pub fn request_from_builder(
274 &mut self,
275 OutRequestBuilder {
276 tid,
277 query_id,
278 peer,
279 command,
280 token,
281 target,
282 value,
283 id,
284 }: OutRequestBuilder,
285 ) -> QueryAndTid {
286 let id = id.or_else(|| (!self.ephemeral).then(|| self.id().0));
287 let tid = tid.unwrap_or_else(|| self.new_tid());
288 self.enqueue_request((
289 query_id,
290 RequestMsgData {
291 tid,
292 to: peer,
293 id,
294 token,
295 command,
296 target: target.map(|x| x.into()),
297 value,
298 },
299 None,
300 ));
301 (query_id, tid)
302 }
303
304 pub fn request(
305 &mut self,
306 command: Command,
307 target: Option<IdBytes>,
308 value: Option<Vec<u8>>,
309 peer: Peer,
310 query_id: Option<QueryId>,
311 token: Option<[u8; 32]>,
312 ) -> QueryAndTid {
313 let id = (!self.ephemeral).then(|| self.id().0);
314 let tid = self.new_tid();
315 self.enqueue_request((
316 query_id,
317 RequestMsgData {
318 tid,
319 to: peer,
320 id,
321 token,
322 command,
323 target: target.map(|x| x.0),
324 value,
325 },
326 None,
327 ));
328 (query_id, tid)
329 }
330
331 pub fn enqueue_request(&mut self, msg: (Option<QueryId>, RequestMsgData, Option<Sender<()>>)) {
332 self.pending_send.push_back(OutMessage::Request(msg));
333 self.maybe_wake();
334 }
335 pub fn error(
336 &mut self,
337 request: &RequestMsgData,
338 error: usize,
339 value: Option<Vec<u8>>,
340 closer_nodes: Option<Vec<Peer>>,
341 peer: &Peer,
342 ) -> crate::Result<()> {
343 let id = (!self.ephemeral).then(|| self.id().0);
344 let token = Some(self.token(peer, 1)?);
345
346 self.enqueue_reply(
347 ReplyMsgData {
348 tid: request.tid,
349 to: peer.clone(),
350 id,
351 token,
352 closer_nodes: closer_nodes.unwrap_or_default(),
353 error,
354 value,
355 },
356 None,
357 );
358 Ok(())
359 }
360
361 pub fn reply(&mut self, mut msg: ReplyMsgData) {
362 if msg.token.is_none() {
363 msg.token = self.token(&msg.to, 1).ok();
364 }
365 self.enqueue_reply(msg, None)
366 }
367
368 pub fn request2(
369 &mut self,
370 OutRequestBuilder {
371 query_id,
372 tid,
373 peer,
374 command,
375 token,
376 target,
377 value,
378 id,
379 }: OutRequestBuilder,
380 ) -> crate::Result<Receiver<()>> {
381 let (tx, rx) = oneshot::channel();
382 let id = id.or_else(|| (!self.ephemeral).then(|| self.id().0));
383 let tid = tid.unwrap_or_else(|| self.new_tid());
384 self.enqueue_request((
385 query_id,
386 RequestMsgData {
387 tid,
388 command,
389 id,
390 token,
391 target: target.map(|t| t.into()),
392 value,
393 to: peer,
394 },
395 Some(tx),
396 ));
397 Ok(rx)
398 }
399
400 pub fn response(
401 &mut self,
402 request: &RequestMsgData,
403 value: Option<Vec<u8>>,
404 closer_nodes: Option<Vec<Peer>>,
405 peer: &Peer,
406 ) -> crate::Result<Receiver<()>> {
407 let id = (!self.ephemeral).then(|| self.id().0);
408 let token = Some(self.token(peer, 1)?);
409 let (tx, rx) = oneshot::channel();
410 self.enqueue_reply(
411 ReplyMsgData {
412 tid: request.tid,
413 to: peer.clone(),
414 id,
415 token,
416 closer_nodes: closer_nodes.unwrap_or_default(),
417 error: 0,
418 value,
419 },
420 Some(tx),
421 );
422 Ok(rx)
423 }
424
425 fn on_response(&mut self, recv: ReplyMsgData, peer: Peer) -> IoHandlerEvent {
426 if let Some(req) = self.pending_recv.remove(&recv.tid) {
427 return IoHandlerEvent::InResponse(Arc::new(InResponse::new(
428 Box::new(req.message),
429 recv,
430 peer,
431 req.query_id,
432 )));
433 }
434 IoHandlerEvent::InResponseBadRequestId {
435 peer,
436 message: recv,
437 }
438 }
439 fn on_message(&mut self, msg: MsgData, rinfo: SocketAddr) -> IoHandlerEvent {
441 let peer = Peer::from(&rinfo);
442 match msg {
443 MsgData::Request(req) => {
444 trace!(name=self.name(), tid = req.tid, command =% req.command, from =? peer.addr, "RX:Request");
445 IoHandlerEvent::InRequest { message: req, peer }
446 }
447 MsgData::Reply(rep) => {
448 trace!(name = self.name(), tid = rep.tid, from =? peer.addr, "RX:Reply");
449 self.on_response(rep, peer)
450 }
451 }
452 }
453
454 fn poll_send(&mut self, cx: &mut Context<'_>) -> Option<IoHandlerEvent> {
455 let msg = match self.pending_flush.take() {
456 Some(m) => m,
457 None => match self.pending_send.pop_front() {
458 Some(m) => m,
459 None => {
460 return match Sink::poll_flush(Pin::new(&mut self.message_stream), cx) {
461 Poll::Ready(_e) => None,
462 Poll::Pending => {
463 cx.waker().wake_by_ref();
464 None
465 }
466 };
467 }
468 },
469 };
470 if !Sink::poll_ready(Pin::new(&mut self.message_stream), cx).is_ready() {
471 self.pending_flush = Some(msg);
472 return None;
473 }
474 let out = match &msg {
475 OutMessage::Request((query_id, message, _tx)) => {
476 let tid = message.tid;
477 self.pending_recv.insert(
478 message.tid,
479 InflightRequest {
480 message: message.clone(),
481 timestamp: Instant::now(),
482 query_id: *query_id,
483 },
484 );
485 IoHandlerEvent::OutRequest { tid }
486 }
487 OutMessage::Reply((_tx, message)) => {
488 let peer = message.to.clone();
489 IoHandlerEvent::OutResponse {
490 message: message.clone(),
491 peer,
492 }
493 }
494 };
495
496 let (msg, socket, tx) = msg.into_sendable();
497 match &msg {
498 MsgData::Request(m) => {
499 trace!(name=self.name(), tid = m.tid, cmd =% m.command, to=?socket, "TX:Request")
500 }
501 MsgData::Reply(m) => trace!(name = self.name(), tid = m.tid, to=?socket, "TX:Reply"),
502 }
503 if let Err(e) = Sink::start_send(Pin::new(&mut self.message_stream), (msg, socket)) {
504 error!(error =? e, "start_send error");
505 todo!()
506 }
507 _ = Sink::poll_flush(Pin::new(&mut self.message_stream), cx);
508 if let Some(tx) = tx {
509 _ = tx.send(());
510 }
511
512 if !self.pending_send.is_empty() {
513 cx.waker().wake_by_ref();
514 }
515 Some(out)
516 }
517
518 fn maybe_wake(&mut self) {
519 if let Some(w) = self.stream_waker.take() {
520 w.wake()
521 }
522 }
523}
524
525#[derive(Debug, Default)]
526pub struct IoConfig {
527 pub secrets: Secrets,
528 pub ephemeral: bool,
531}
532
533impl Stream for IoHandler {
534 type Item = IoHandlerEvent;
535
536 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
537 let pin = self.get_mut();
538 _ = pin.stream_waker.insert(cx.waker().clone());
539
540 if let Some(out) = pin.poll_send(cx) {
541 return Poll::Ready(Some(out));
542 }
543
544 match Stream::poll_next(Pin::new(&mut pin.message_stream), cx) {
546 Poll::Ready(Some(Ok((msg, rinfo)))) => {
547 let out = pin.on_message(msg, rinfo);
548 cx.waker().wake_by_ref();
549 return Poll::Ready(Some(out));
550 }
551 Poll::Ready(Some(Err(err))) => {
552 let out = IoHandlerEvent::InSocketErr { err };
553 error!(name = pin.name(), "{out:#?}");
554 return Poll::Ready(Some(out));
555 }
556 _ => {}
557 }
558
559 Poll::Pending
564 }
565}
566
567#[derive(Debug)]
569pub enum IoHandlerEvent {
570 OutResponse { message: ReplyMsgData, peer: Peer },
572 OutRequest { tid: Tid },
574 InResponse(Arc<InResponse>),
576 InRequest { message: RequestMsgData, peer: Peer },
578 OutSocketErr { err: crate::Error },
580 RequestTimeout {
582 message: MsgData,
583 peer: Peer,
584 sent: Instant,
585 query_id: QueryId,
586 },
587 InMessageErr { err: io::Error, peer: Peer },
590 InSocketErr { err: crate::Error },
592 InResponseBadRequestId { message: ReplyMsgData, peer: Peer },
595 ChanneledResponse(Tid),
597}
598
599impl IoHandlerEvent {
600 fn kind(&self) -> String {
601 use IoHandlerEvent as Ihe;
602 match self {
603 Ihe::OutResponse { .. } => "OutResponse",
604 Ihe::OutRequest { .. } => "OutRequest",
605 Ihe::InResponse(_) => "InResponse",
606 Ihe::InRequest { .. } => "InRequest",
607 Ihe::OutSocketErr { .. } => "OutSocketErr",
608 Ihe::RequestTimeout { .. } => "RequestTimeout",
609 Ihe::InMessageErr { .. } => "InMessageErr",
610 Ihe::InSocketErr { .. } => "InSocketErr",
611 Ihe::InResponseBadRequestId { .. } => "InResponseBadRequestId",
612 Ihe::ChanneledResponse(_) => "ChanneledResponse",
613 }
614 .to_string()
615 }
616}
617
618impl std::fmt::Display for IoHandlerEvent {
619 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
620 use IoHandlerEvent as Ihe;
621 match self {
622 Ihe::InResponse(x) => write!(
623 f,
624 "InRespInResponse(tid={}, cmd={})",
625 x.request.tid, x.request.command
626 ),
627 _ => write!(f, "{}()", self.kind()),
628 }
629 }
630}
631
632pub fn random_name() -> String {
634 use rand::Rng;
635 const CHARSET: &[u8] = b"abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ";
636 let mut rng = rand::thread_rng();
637 (0..5)
638 .map(|_| {
639 let idx = rng.gen_range(0, CHARSET.len());
640 CHARSET[idx] as char
641 })
642 .collect()
643}
644
645#[cfg(test)]
646mod test {
647 use crate::{InternalCommand, stateobserver::State, thirty_two_random_bytes};
648 use futures::StreamExt;
649
650 use super::*;
651
652 fn new_io() -> IoHandler {
653 let view = State::new(IdBytes::from(thirty_two_random_bytes())).view();
654 let message_stream = MessageDataStream::defualt_bind().unwrap();
655 IoHandler::new(view, message_stream, Default::default())
656 }
657 #[tokio::test]
658 async fn test_iohandler_to_iohandler_messaging() -> crate::Result<()> {
659 let mut a = new_io();
660 let mut b = new_io();
661
662 let to = Peer::from(&b.local_addr()?);
663 let id = Some(thirty_two_random_bytes());
664 let msg = RequestMsgData {
665 tid: 42,
666 to,
667 id,
668 token: None,
669 command: InternalCommand::Ping.into(),
670 target: None,
671 value: None,
672 };
673 let query_id = Some(QueryId(42));
674 a.enqueue_request((query_id, msg.clone(), None));
675 a.next().await;
676 let IoHandlerEvent::InRequest { message: res, .. } = b.next().await.unwrap() else {
677 panic!()
678 };
679 assert_eq!(res, msg);
680 Ok(())
681 }
682}