node: Disconnect peer for wrong version or magic

Signed-off-by: Alexis Sellier <self@cloudhead.io>
This commit is contained in:
Alexis Sellier 2022-08-23 13:28:21 +02:00
parent 38f4bb52d8
commit e078706cd1
No known key found for this signature in database
1 changed files with 33 additions and 9 deletions

View File

@ -200,10 +200,9 @@ impl<T: ReadStorage + WriteStorage, S: address_book::Store> Protocol<S, T> {
} }
} }
pub fn disconnect(&mut self, peer: &PeerId) { pub fn disconnect(&mut self, peer: &PeerId, reason: DisconnectReason) {
if let Some(addr) = self.peers.get(peer).map(|p| p.addr) { if let Some(addr) = self.peers.get(peer).map(|p| p.addr) {
self.outbox() self.context.disconnect(addr, reason);
.push_back(Io::Disconnect(addr, DisconnectReason::User));
} }
} }
@ -533,7 +532,10 @@ where
}; };
for msg in msgs { for msg in msgs {
peer.received(msg, &mut self.context); if let Err(err) = peer.received(msg, &mut self.context) {
self.context
.disconnect(peer.addr, DisconnectReason::Error(err));
}
} }
} }
} }
@ -555,12 +557,14 @@ impl<S, T> DerefMut for Protocol<S, T> {
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub enum DisconnectReason { pub enum DisconnectReason {
User, User,
Error(PeerError),
} }
impl DisconnectReason { impl DisconnectReason {
fn is_transient(&self) -> bool { fn is_transient(&self) -> bool {
match self { match self {
Self::User => false, Self::User => false,
Self::Error(..) => false,
} }
} }
} }
@ -575,6 +579,7 @@ impl fmt::Display for DisconnectReason {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self { match self {
Self::User => write!(f, "user"), Self::User => write!(f, "user"),
Self::Error(err) => write!(f, "error: {}", err),
} }
} }
} }
@ -654,6 +659,11 @@ where
inventory.insert(from); inventory.insert(from);
} }
} }
/// Disconnect a peer.
fn disconnect(&mut self, addr: net::SocketAddr, reason: DisconnectReason) {
self.io.push_back(Io::Disconnect(addr, reason));
}
} }
impl<S, T> Context<S, T> { impl<S, T> Context<S, T> {
@ -725,6 +735,14 @@ enum PeerState {
Disconnected { since: LocalTime }, Disconnected { since: LocalTime },
} }
#[derive(thiserror::Error, Debug, Clone)]
pub enum PeerError {
#[error("wrong network constant in message: {0}")]
WrongMagic(u32),
#[error("wrong protocol version in message: {0}")]
WrongVersion(u32),
}
#[derive(Debug)] #[derive(Debug)]
pub struct Peer { pub struct Peer {
/// Peer address. /// Peer address.
@ -764,20 +782,25 @@ impl Peer {
matches!(self.state, PeerState::Negotiated { .. }) matches!(self.state, PeerState::Negotiated { .. })
} }
fn received<S, T>(&mut self, envelope: Envelope, ctx: &mut Context<S, T>) fn received<S, T>(
&mut self,
envelope: Envelope,
ctx: &mut Context<S, T>,
) -> Result<(), PeerError>
where where
T: storage::ReadStorage, T: storage::ReadStorage,
{ {
if envelope.magic != NETWORK_MAGIC { if envelope.magic != NETWORK_MAGIC {
// TODO: Disconnect return Err(PeerError::WrongMagic(envelope.magic));
return;
} }
debug!("Received {:?} from {}", &envelope.msg, self.id()); debug!("Received {:?} from {}", &envelope.msg, self.id());
match envelope.msg { match envelope.msg {
Message::Hello { .. } => { Message::Hello { version } => {
if version != PROTOCOL_VERSION {
return Err(PeerError::WrongVersion(version));
}
if let PeerState::Initial = self.state { if let PeerState::Initial = self.state {
// TODO: Check version.
// Nb. This is a very primitive handshake. Eventually we should have anyhow // Nb. This is a very primitive handshake. Eventually we should have anyhow
// extra "acknowledgment" message sent when the `Hello` is well received. // extra "acknowledgment" message sent when the `Hello` is well received.
if self.link.is_inbound() { if self.link.is_inbound() {
@ -807,5 +830,6 @@ impl Peer {
todo!(); todo!();
} }
} }
Ok(())
} }
} }