Skip to main content

dht_rpc/
io.rs

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/// TODO hide secrets in fmt::Debug
36#[derive(Debug)]
37pub struct Secrets {
38    _rotate_counter: usize,
39    // NB starts null in js. Not initialized until token call.
40    // so my behavior diverges when drain called until token
41    // bc drain checks if secrets initialized
42    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/// Recied response data along with metadata
84#[derive(Debug, Clone)]
85pub struct InResponse {
86    pub request: Box<RequestMsgData>,
87    pub response: ReplyMsgData,
88    /// [`Peer`] who sent the response
89    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/// OutMessage contains outgoing messages data, including local metadata for managing messages
177#[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    /// The message send
201    message: RequestMsgData,
202    /// Timestamp when the request was sent
203    #[expect(unused)] // TODO FIXME not read. Why not?
204    timestamp: Instant,
205    // Identifier for the query this request is used with
206    query_id: Option<QueryId>,
207}
208
209#[derive(Debug)]
210pub struct IoHandler {
211    id: Observer<IdBytes>,
212    ephemeral: bool,
213    message_stream: MessageDataStream,
214    /// Messages to send
215    pending_send: VecDeque<OutMessage>,
216    /// Current message
217    pending_flush: Option<OutMessage>,
218    /// Sent requests we currently wait for a response
219    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    /// TODO check this is correct.
260    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    /// A new `Message` was read from the socket.
440    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    /// When true, the node won't expose its ID to remote peers.
529    /// Defaults to false (non-ephemeral).
530    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        // read from socket
545        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        //if pin.last_rotation + pin.rotation > Instant::now() {
560        //    pin.rotate_secrets();
561        //}
562
563        Poll::Pending
564    }
565}
566
567/// Event generated by the IO handler
568#[derive(Debug)]
569pub enum IoHandlerEvent {
570    ///  A response was sent
571    OutResponse { message: ReplyMsgData, peer: Peer },
572    /// A request was sent
573    OutRequest { tid: Tid },
574    /// A Response to a Query Message was recieved
575    InResponse(Arc<InResponse>),
576    /// A Request was receieved
577    InRequest { message: RequestMsgData, peer: Peer },
578    /// Error while sending a message
579    OutSocketErr { err: crate::Error },
580    /// A request did not recieve a response within the given timeout
581    RequestTimeout {
582        message: MsgData,
583        peer: Peer,
584        sent: Instant,
585        query_id: QueryId,
586    },
587    /// Error while decoding a message from socket
588    /// TODO unused
589    InMessageErr { err: io::Error, peer: Peer },
590    /// Error while reading from socket
591    InSocketErr { err: crate::Error },
592    /// Received a response with a request id that was doesn't match any pending
593    /// responses.
594    InResponseBadRequestId { message: ReplyMsgData, peer: Peer },
595    /// A Response to message handled by a request future
596    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
632/// return a Random String, 5 letters long containing only the letters a-zA-Z.
633pub 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}