From 46536509d70ea304120a206d183d41653d29731c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Andr=C3=A9=20Colomb?= Date: Sun, 7 Jun 2020 10:31:12 +0200 Subject: [PATCH] lib/protocol: Avoid panic in DeviceIDFromBytes (#6714) --- cmd/stindex/dump.go | 14 +++++++++----- cmd/strelaysrv/listener.go | 10 +++++++++- lib/db/db_test.go | 6 +++++- lib/db/lowlevel.go | 5 ++++- lib/db/meta.go | 11 +++++++++-- lib/db/transactions.go | 21 +++++++++++++++++---- lib/protocol/deviceid.go | 6 +++--- lib/protocol/deviceid_test.go | 11 ++++++++--- lib/relay/protocol/packets.go | 6 +++++- 9 files changed, 69 insertions(+), 21 deletions(-) diff --git a/cmd/stindex/dump.go b/cmd/stindex/dump.go index 9e989a492..ad4faac6d 100644 --- a/cmd/stindex/dump.go +++ b/cmd/stindex/dump.go @@ -73,12 +73,16 @@ func dump(ldb backend.Backend) { case db.KeyTypeDeviceIdx: key := binary.BigEndian.Uint32(key[1:]) val := it.Value() - if len(val) == 0 { - fmt.Printf("[deviceidx] K:%d V:\n", key) - } else { - dev := protocol.DeviceIDFromBytes(val) - fmt.Printf("[deviceidx] K:%d V:%s\n", key, dev) + device := "" + if len(val) > 0 { + dev, err := protocol.DeviceIDFromBytes(val) + if err != nil { + device = fmt.Sprintf("", len(val)) + } else { + device = dev.String() + } } + fmt.Printf("[deviceidx] K:%d V:%s\n", key, device) case db.KeyTypeIndexID: device := binary.BigEndian.Uint32(key[1:]) diff --git a/cmd/strelaysrv/listener.go b/cmd/strelaysrv/listener.go index 6ce3fc4bb..c0e033c2d 100644 --- a/cmd/strelaysrv/listener.go +++ b/cmd/strelaysrv/listener.go @@ -149,7 +149,15 @@ func protocolConnectionHandler(tcpConn net.Conn, config *tls.Config) { protocol.WriteMessage(conn, protocol.ResponseSuccess) case protocol.ConnectRequest: - requestedPeer := syncthingprotocol.DeviceIDFromBytes(msg.ID) + requestedPeer, err := syncthingprotocol.DeviceIDFromBytes(msg.ID) + if err != nil { + if debug { + log.Println(id, "is looking for an invalid peer ID") + } + protocol.WriteMessage(conn, protocol.ResponseNotFound) + conn.Close() + continue + } outboxesMut.RLock() peerOutbox, ok := outboxes[requestedPeer] outboxesMut.RUnlock() diff --git a/lib/db/db_test.go b/lib/db/db_test.go index b16459edb..28d6efb61 100644 --- a/lib/db/db_test.go +++ b/lib/db/db_test.go @@ -264,7 +264,11 @@ func TestUpdate0to3(t *testing.T) { t.Fatal(err) } if !ok { - t.Fatal("surprise missing global file", string(name), protocol.DeviceIDFromBytes(vl.Versions[0].Device)) + device := "" + if dev, err := protocol.DeviceIDFromBytes(vl.Versions[0].Device); err != nil { + device = dev.String() + } + t.Fatal("surprise missing global file", string(name), device) } e, ok := need[fi.FileName()] if !ok { diff --git a/lib/db/lowlevel.go b/lib/db/lowlevel.go index d3ddc4680..fe3ec685a 100644 --- a/lib/db/lowlevel.go +++ b/lib/db/lowlevel.go @@ -121,7 +121,10 @@ func (db *Lowlevel) updateRemoteFiles(folder, device []byte, fs []protocol.FileI defer t.close() var dk, gk, keyBuf []byte - devID := protocol.DeviceIDFromBytes(device) + devID, err := protocol.DeviceIDFromBytes(device) + if err != nil { + return err + } for _, f := range fs { name := []byte(f.Name) dk, err = db.keyer.GenerateDeviceFileKey(dk, folder, device, name) diff --git a/lib/db/meta.go b/lib/db/meta.go index baa376aed..5f682898f 100644 --- a/lib/db/meta.go +++ b/lib/db/meta.go @@ -52,7 +52,11 @@ func (m *metadataTracker) Unmarshal(bs []byte) error { // Initialize the index map for i, c := range m.counts.Counts { - m.indexes[metaKey{protocol.DeviceIDFromBytes(c.DeviceID), c.LocalFlags}] = i + dev, err := protocol.DeviceIDFromBytes(c.DeviceID) + if err != nil { + return err + } + m.indexes[metaKey{dev, c.LocalFlags}] = i } return nil } @@ -392,7 +396,10 @@ func (m *countsMap) devices() []protocol.DeviceID { for _, dev := range m.counts.Counts { if dev.Sequence > 0 { - id := protocol.DeviceIDFromBytes(dev.DeviceID) + id, err := protocol.DeviceIDFromBytes(dev.DeviceID) + if err != nil { + panic(err) + } if id == protocol.GlobalDeviceID || id == protocol.LocalDeviceID { continue } diff --git a/lib/db/transactions.go b/lib/db/transactions.go index 05cacbaac..bc55087c9 100644 --- a/lib/db/transactions.go +++ b/lib/db/transactions.go @@ -414,7 +414,11 @@ func (t *readOnlyTransaction) availability(folder, file []byte) ([]protocol.Devi } devices := make([]protocol.DeviceID, len(fv.Devices)) for i, dev := range fv.Devices { - devices[i] = protocol.DeviceIDFromBytes(dev) + n, err := protocol.DeviceIDFromBytes(dev) + if err != nil { + return nil, err + } + devices[i] = n } return devices, nil @@ -436,7 +440,10 @@ func (t *readOnlyTransaction) withNeed(folder, device []byte, truncate bool, fn defer dbi.Release() var dk []byte - devID := protocol.DeviceIDFromBytes(device) + devID, err := protocol.DeviceIDFromBytes(device) + if err != nil { + return err + } for dbi.Next() { var vl VersionList if err := vl.Unmarshal(dbi.Value()); err != nil { @@ -592,7 +599,10 @@ func (t readWriteTransaction) putFile(fkey []byte, fi protocol.FileInfo, truncat // file. If the device is already present in the list, the version is updated. // If the file does not have an entry in the global list, it is created. func (t readWriteTransaction) updateGlobal(gk, keyBuf, folder, device []byte, file protocol.FileInfo, meta *metadataTracker) ([]byte, bool, error) { - deviceID := protocol.DeviceIDFromBytes(device) + deviceID, err := protocol.DeviceIDFromBytes(device) + if err != nil { + return nil, false, err + } l.Debugf("update global; folder=%q device=%v file=%q version=%v invalid=%v", folder, deviceID, file.Name, file.Version, file.IsInvalid()) @@ -768,7 +778,10 @@ func need(global FileVersion, haveLocal bool, localVersion protocol.Vector) bool // given file. If the version list is empty after this, the file entry is // removed entirely. func (t readWriteTransaction) removeFromGlobal(gk, keyBuf, folder, device, file []byte, meta *metadataTracker) ([]byte, error) { - deviceID := protocol.DeviceIDFromBytes(device) + deviceID, err := protocol.DeviceIDFromBytes(device) + if err != nil { + return nil, err + } l.Debugf("remove from global; folder=%q device=%v file=%q", folder, deviceID, file) diff --git a/lib/protocol/deviceid.go b/lib/protocol/deviceid.go index 8cac191ca..a721d44e2 100644 --- a/lib/protocol/deviceid.go +++ b/lib/protocol/deviceid.go @@ -46,13 +46,13 @@ func DeviceIDFromString(s string) (DeviceID, error) { return n, err } -func DeviceIDFromBytes(bs []byte) DeviceID { +func DeviceIDFromBytes(bs []byte) (DeviceID, error) { var n DeviceID if len(bs) != len(n) { - panic("incorrect length of byte slice representing device ID") + return n, fmt.Errorf("incorrect length of byte slice representing device ID") } copy(n[:], bs) - return n + return n, nil } // String returns the canonical string representation of the device ID diff --git a/lib/protocol/deviceid_test.go b/lib/protocol/deviceid_test.go index d7c05313a..a11e8e2c1 100644 --- a/lib/protocol/deviceid_test.go +++ b/lib/protocol/deviceid_test.go @@ -99,8 +99,10 @@ func TestShortIDString(t *testing.T) { func TestDeviceIDFromBytes(t *testing.T) { id0, _ := DeviceIDFromString(formatted) - id1 := DeviceIDFromBytes(id0[:]) - if id1.String() != formatted { + id1, err := DeviceIDFromBytes(id0[:]) + if err != nil { + t.Fatal(err) + } else if id1.String() != formatted { t.Errorf("Wrong device ID, got %q, want %q", id1, formatted) } } @@ -150,7 +152,10 @@ func TestNewDeviceIDMarshalling(t *testing.T) { // Verify it's the same - if DeviceIDFromBytes(msg2.Test) != id0 { + id1, err := DeviceIDFromBytes(msg2.Test) + if err != nil { + t.Fatal(err) + } else if id1 != id0 { t.Error("Mismatch in old -> new direction") } } diff --git a/lib/relay/protocol/packets.go b/lib/relay/protocol/packets.go index 9c329ec00..ea790883a 100644 --- a/lib/relay/protocol/packets.go +++ b/lib/relay/protocol/packets.go @@ -56,7 +56,11 @@ type SessionInvitation struct { } func (i SessionInvitation) String() string { - return fmt.Sprintf("%s@%s", syncthingprotocol.DeviceIDFromBytes(i.From), i.AddressString()) + device := "" + if address, err := syncthingprotocol.DeviceIDFromBytes(i.From); err == nil { + device = address.String() + } + return fmt.Sprintf("%s@%s", device, i.AddressString()) } func (i SessionInvitation) GoString() string {