mirror of
https://github.com/ethereum/go-ethereum.git
synced 2026-08-18 09:53:48 +00:00
p2p: handle both disconnect encodings
And update some unit tests, since they were still using the wrong format.
This commit is contained in:
parent
9950b419ad
commit
d109446f86
4 changed files with 36 additions and 11 deletions
25
p2p/peer.go
25
p2p/peer.go
|
|
@ -345,9 +345,7 @@ func (p *Peer) handle(msg Msg) error {
|
|||
case msg.Code == discMsg:
|
||||
// This is the last message. We don't need to discard or
|
||||
// check errors because, the connection will be closed after it.
|
||||
var m struct{ R DiscReason }
|
||||
rlp.Decode(msg.Payload, &m)
|
||||
return m.R
|
||||
return decodeDisconnectMessage(msg.Payload)
|
||||
case msg.Code < baseProtocolLength:
|
||||
// ignore other base protocol messages
|
||||
return msg.Discard()
|
||||
|
|
@ -372,6 +370,27 @@ func (p *Peer) handle(msg Msg) error {
|
|||
return nil
|
||||
}
|
||||
|
||||
// decodeDisconnectMessage decodes the payload of discMsg.
|
||||
func decodeDisconnectMessage(r io.Reader) (reason DiscReason) {
|
||||
s := rlp.NewStream(r, 100)
|
||||
k, _, err := s.Kind()
|
||||
if err != nil {
|
||||
return DiscInvalid
|
||||
}
|
||||
if k == rlp.List {
|
||||
s.List()
|
||||
err = s.Decode(&reason)
|
||||
} else {
|
||||
// Legacy path: some implementations, including geth, used to send the disconnect
|
||||
// reason as a byte array by accident.
|
||||
err = s.Decode(&reason)
|
||||
}
|
||||
if err != nil {
|
||||
reason = DiscInvalid
|
||||
}
|
||||
return reason
|
||||
}
|
||||
|
||||
func countMatchingProtocols(protocols []Protocol, caps []Cap) int {
|
||||
n := 0
|
||||
for _, cap := range caps {
|
||||
|
|
|
|||
|
|
@ -70,6 +70,8 @@ const (
|
|||
DiscSelf
|
||||
DiscReadTimeout
|
||||
DiscSubprotocolError = DiscReason(0x10)
|
||||
|
||||
DiscInvalid = 0xff
|
||||
)
|
||||
|
||||
var discReasonToString = [...]string{
|
||||
|
|
@ -86,6 +88,7 @@ var discReasonToString = [...]string{
|
|||
DiscSelf: "connected to self",
|
||||
DiscReadTimeout: "read timeout",
|
||||
DiscSubprotocolError: "subprotocol error",
|
||||
DiscInvalid: "invalid disconnect reason",
|
||||
}
|
||||
|
||||
func (d DiscReason) String() string {
|
||||
|
|
|
|||
|
|
@ -120,7 +120,7 @@ func (t *rlpxTransport) close(err error) {
|
|||
if err := t.conn.SetWriteDeadline(deadline); err == nil {
|
||||
// Connection supports write deadline.
|
||||
t.wbuf.Reset()
|
||||
rlp.Encode(&t.wbuf, []interface{}{reason})
|
||||
rlp.Encode(&t.wbuf, []any{reason})
|
||||
t.conn.Write(discMsg, t.wbuf.Bytes())
|
||||
}
|
||||
}
|
||||
|
|
@ -164,11 +164,8 @@ func readProtocolHandshake(rw MsgReader) (*protoHandshake, error) {
|
|||
if msg.Code == discMsg {
|
||||
// Disconnect before protocol handshake is valid according to the
|
||||
// spec and we send it ourself if the post-handshake checks fail.
|
||||
// We can't return the reason directly, though, because it is echoed
|
||||
// back otherwise. Wrap it in a string instead.
|
||||
var m struct{ R DiscReason }
|
||||
rlp.Decode(msg.Payload, &m)
|
||||
return nil, m.R
|
||||
r := decodeDisconnectMessage(msg.Payload)
|
||||
return nil, r
|
||||
}
|
||||
if msg.Code != handshakeMsg {
|
||||
return nil, fmt.Errorf("expected handshake, got %x", msg.Code)
|
||||
|
|
|
|||
|
|
@ -97,7 +97,7 @@ func TestProtocolHandshake(t *testing.T) {
|
|||
return
|
||||
}
|
||||
|
||||
if err := ExpectMsg(rlpx, discMsg, []DiscReason{DiscQuitting}); err != nil {
|
||||
if err := ExpectMsg(rlpx, discMsg, []any{DiscQuitting}); err != nil {
|
||||
t.Errorf("error receiving disconnect: %v", err)
|
||||
}
|
||||
}()
|
||||
|
|
@ -112,7 +112,13 @@ func TestProtocolHandshakeErrors(t *testing.T) {
|
|||
}{
|
||||
{
|
||||
code: discMsg,
|
||||
msg: []DiscReason{DiscQuitting},
|
||||
msg: []any{DiscQuitting},
|
||||
err: DiscQuitting,
|
||||
},
|
||||
{
|
||||
// legacy disconnect encoding as byte array
|
||||
code: discMsg,
|
||||
msg: []byte{byte(DiscQuitting)},
|
||||
err: DiscQuitting,
|
||||
},
|
||||
{
|
||||
|
|
|
|||
Loading…
Reference in a new issue