diff --git a/daemon/libnetwork/drivers/overlay/encryption.go b/daemon/libnetwork/drivers/overlay/encryption.go index 8a436b4d79..258ae3b92d 100644 --- a/daemon/libnetwork/drivers/overlay/encryption.go +++ b/daemon/libnetwork/drivers/overlay/encryption.go @@ -135,7 +135,7 @@ func (d *driver) setupEncryption(remoteIP netip.Addr) error { indices := make([]spi, 0, len(d.keys)) for i, k := range d.keys { - spis := spi{buildSPI(advIP.AsSlice(), remoteIP.AsSlice(), k.tag), buildSPI(remoteIP.AsSlice(), advIP.AsSlice(), k.tag)} + spis := spi{buildSPI(advIP, remoteIP, k.tag), buildSPI(remoteIP, advIP, k.tag)} dir := reverse if i == 0 { dir = bidir @@ -430,14 +430,14 @@ func spExists(sp *netlink.XfrmPolicy) (bool, error) { } } -func buildSPI(src, dst net.IP, st uint32) int { - b := make([]byte, 4) - binary.BigEndian.PutUint32(b, st) +func buildSPI(src, dst netip.Addr, st uint32) int { h := fnv.New32a() - h.Write(src) - h.Write(b) - h.Write(dst) - return int(binary.BigEndian.Uint32(h.Sum(nil))) + v := src.As16() + h.Write(v[:]) + binary.Write(h, binary.BigEndian, st) + v = dst.As16() + h.Write(v[:]) + return int(h.Sum32()) } func buildAeadAlgo(k *key, s int) *netlink.XfrmStateAlgo { @@ -542,7 +542,7 @@ func (d *driver) updateKeys(ctx context.Context, encrData discoverapi.DriverEncr } for rIP, node := range d.secMap { - idxs := updateNodeKey(lIP.AsSlice(), aIP.AsSlice(), rIP.AsSlice(), node.spi, d.keys, newIdx, priIdx, delIdx) + idxs := updateNodeKey(lIP, aIP, rIP, node.spi, d.keys, newIdx, priIdx, delIdx) if idxs != nil { d.secMap[rIP] = encrNode{idxs, node.count} } @@ -574,7 +574,7 @@ func (d *driver) updateKeys(ctx context.Context, encrData discoverapi.DriverEncr *********************************************************/ // Spis and keys are sorted in such away the one in position 0 is the primary -func updateNodeKey(lIP, aIP, rIP net.IP, idxs []spi, curKeys []*key, newIdx, priIdx, delIdx int) []spi { +func updateNodeKey(lIP, aIP, rIP netip.Addr, idxs []spi, curKeys []*key, newIdx, priIdx, delIdx int) []spi { log.G(context.TODO()).Debugf("Updating keys for node: %s (%d,%d,%d)", rIP, newIdx, priIdx, delIdx) spis := idxs @@ -590,17 +590,17 @@ func updateNodeKey(lIP, aIP, rIP net.IP, idxs []spi, curKeys []*key, newIdx, pri if delIdx != -1 { // -rSA0 - programSA(lIP, rIP, spis[delIdx], nil, reverse, false) + programSA(lIP.AsSlice(), rIP.AsSlice(), spis[delIdx], nil, reverse, false) } if newIdx > -1 { // +rSA2 - programSA(lIP, rIP, spis[newIdx], curKeys[newIdx], reverse, true) + programSA(lIP.AsSlice(), rIP.AsSlice(), spis[newIdx], curKeys[newIdx], reverse, true) } if priIdx > 0 { // +fSA2 - fSA2, _, _ := programSA(lIP, rIP, spis[priIdx], curKeys[priIdx], forward, true) + fSA2, _, _ := programSA(lIP.AsSlice(), rIP.AsSlice(), spis[priIdx], curKeys[priIdx], forward, true) // +fSP2, -fSP1 s := getMinimalIP(fSA2.Src) @@ -631,7 +631,7 @@ func updateNodeKey(lIP, aIP, rIP net.IP, idxs []spi, curKeys []*key, newIdx, pri } // -fSA1 - programSA(lIP, rIP, spis[0], nil, forward, false) + programSA(lIP.AsSlice(), rIP.AsSlice(), spis[0], nil, forward, false) } // swap diff --git a/daemon/libnetwork/drivers/overlay/encryption_test.go b/daemon/libnetwork/drivers/overlay/encryption_test.go new file mode 100644 index 0000000000..8df98eb65e --- /dev/null +++ b/daemon/libnetwork/drivers/overlay/encryption_test.go @@ -0,0 +1,51 @@ +//go:build linux + +package overlay + +import ( + "encoding/binary" + "hash/fnv" + "net" + "net/netip" + "testing" +) + +func legacyBuildSPI(src, dst net.IP, st uint32) int { + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, st) + h := fnv.New32a() + h.Write(src) + h.Write(b) + h.Write(dst) + return int(binary.BigEndian.Uint32(h.Sum(nil))) +} + +func TestBuildSPI(t *testing.T) { + cases := []struct { + src, dst string + st uint32 + }{ + {"1.2.3.4", "5.6.7.8", 1234}, + {"::ffff:1.2.3.4", "::ffff:5.6.7.8", 1234}, + {"10.0.0.1", "2001:db8::1", 5678}, + {"2002::abcd:1", "172.15.14.13", 54321}, + {"2002:db8::42", "2002:db8::69", 9999}, + } + for _, tc := range cases { + // The legacy buildSPI function is sensitive to whether src and + // dst are in 4-byte or 16-byte form. Versions of the driver + // using this function always pass the results of net.ParseIP(), + // which always parses the address into 16-byte form. + expected := legacyBuildSPI(net.ParseIP(tc.src), net.ParseIP(tc.dst), tc.st) + + src, dst := netip.MustParseAddr(tc.src), netip.MustParseAddr(tc.dst) + actual := buildSPI(src, dst, tc.st) + if expected != actual { + t.Errorf("buildSPI(%v, %v, %v) = %v; want %v", src, dst, tc.st, actual, expected) + } + actual = buildSPI(src.Unmap(), dst.Unmap(), tc.st) + if expected != actual { + t.Errorf("buildSPI(%v, %v, %v) = %v; want %v", src.Unmap(), dst.Unmap(), tc.st, actual, expected) + } + } +}