go-libp2p: CWE-407 matchMuxers O(I*R) slice scan in Noise handshake, MOAD-0002 through 0005 CLEAN

This commit is contained in:
russell@unturf.com 2026-03-31 13:18:16 -04:00
parent 831e13b1d6
commit 3001d5fcf9
2 changed files with 178 additions and 0 deletions

View file

@ -0,0 +1,29 @@
--- a/p2p/security/noise/transport.go
+++ b/p2p/security/noise/transport.go
@@ -3,7 +3,6 @@ package noise
import (
"context"
"net"
- "slices"
"github.com/libp2p/go-libp2p/core/crypto"
"github.com/libp2p/go-libp2p/core/peer"
@@ -88,10 +87,13 @@ func (i *transportEarlyDataHandler) Received(_ context.Context, _ net.Conn, ext
func matchMuxers(initiatorMuxers, responderMuxers []protocol.ID) protocol.ID {
- for _, initMuxer := range initiatorMuxers {
- if slices.Contains(responderMuxers, initMuxer) {
- return initMuxer
+ if len(initiatorMuxers) == 0 || len(responderMuxers) == 0 {
+ return ""
+ }
+ respSet := make(map[protocol.ID]struct{}, len(responderMuxers))
+ for _, m := range responderMuxers {
+ respSet[m] = struct{}{}
+ }
+ for _, initMuxer := range initiatorMuxers {
+ if _, ok := respSet[initMuxer]; ok {
+ return initMuxer
}
}
return ""

View file

@ -0,0 +1,149 @@
#!/usr/bin/env python3
"""
Unit test for go-libp2p-0002: matchMuxers() uses slices.Contains inside
a loop = O(I*R) where I = initiator muxers, R = responder muxers.
The remote peer controls the size of responderMuxers (up to maxProtoNum=100).
Fix: build a map[protocol.ID]struct{} from responderMuxers first = O(R)
then loop initiatorMuxers with O(1) lookups = O(I+R) total.
File: p2p/security/noise/transport.go
Function: matchMuxers(initiatorMuxers, responderMuxers []protocol.ID) protocol.ID
"""
import time
import unittest
# --- DEFECTIVE VERSION: O(I * R) ---
def match_muxers_defective(initiator_muxers, responder_muxers):
"""Original: linear Contains inside outer loop."""
for init_muxer in initiator_muxers:
if init_muxer in responder_muxers: # O(R) linear scan
return init_muxer
return ""
# --- PATCHED VERSION: O(I + R) ---
def match_muxers_patched(initiator_muxers, responder_muxers):
"""Patched: build set from responder list, then O(1) per initiator."""
if not responder_muxers or not initiator_muxers:
return ""
responder_set = set(responder_muxers)
for init_muxer in initiator_muxers:
if init_muxer in responder_set:
return init_muxer
return ""
class TestMatchMuxersCorrectness(unittest.TestCase):
"""Verify patched output matches defective output."""
def _check(self, init_muxers, resp_muxers, expected):
result_d = match_muxers_defective(init_muxers, resp_muxers)
result_p = match_muxers_patched(init_muxers, resp_muxers)
self.assertEqual(result_d, result_p, "results differ between versions")
self.assertEqual(result_p, expected)
def test_first_match_preferred(self):
self._check(
["/yamux/1.0.0", "/mplex/6.7.0"],
["/mplex/6.7.0", "/yamux/1.0.0"],
"/yamux/1.0.0"
)
def test_only_second_matches(self):
self._check(
["/yamux/1.0.0", "/mplex/6.7.0"],
["/mplex/6.7.0"],
"/mplex/6.7.0"
)
def test_no_match(self):
self._check(
["/yamux/1.0.0"],
["/mplex/6.7.0"],
""
)
def test_empty_initiator(self):
self._check([], ["/yamux/1.0.0"], "")
def test_empty_responder(self):
self._check(["/yamux/1.0.0"], [], "")
def test_both_empty(self):
self._check([], [], "")
def test_single_match(self):
self._check(
["/yamux/1.0.0"],
["/yamux/1.0.0"],
"/yamux/1.0.0"
)
class TestMatchMuxersPerformance(unittest.TestCase):
"""Benchmark: at maxProtoNum=100 the patched version should be faster."""
def test_performance_at_scale(self):
# Remote peer sends maxProtoNum=100 muxers. We send 10 muxers.
# Worst case: no match found, all (I * R) = 1000 comparisons for defective.
# Patched: O(I + R) = 110 operations.
max_proto = 100
resp_muxers = [f"/fake-muxer/{i}" for i in range(max_proto)]
init_muxers = [f"/init-muxer/{i}" for i in range(10)] # 10 initiator muxers
iters = 2000
start = time.perf_counter()
for _ in range(iters):
match_muxers_defective(init_muxers, resp_muxers)
defective_time = time.perf_counter() - start
start = time.perf_counter()
for _ in range(iters):
match_muxers_patched(init_muxers, resp_muxers)
patched_time = time.perf_counter() - start
ratio = defective_time / patched_time if patched_time > 0 else float('inf')
print(f"\n I=10 init, R=100 resp, {iters} iters:")
print(f" defective={defective_time:.4f}s patched={patched_time:.4f}s ratio={ratio:.1f}x")
self.assertGreater(
ratio, 2.0,
f"expected patched to be at least 2x faster at I=10, R=100, got {ratio:.1f}x"
)
def test_performance_large_initiator(self):
# Both sides send maxProtoNum=100 muxers with no match (worst case O(I*R)).
max_proto = 100
init_muxers = [f"/init-muxer/{i}" for i in range(max_proto)]
resp_muxers = [f"/resp-muxer/{i}" for i in range(max_proto)]
iters = 500
start = time.perf_counter()
for _ in range(iters):
match_muxers_defective(init_muxers, resp_muxers)
defective_time = time.perf_counter() - start
start = time.perf_counter()
for _ in range(iters):
match_muxers_patched(init_muxers, resp_muxers)
patched_time = time.perf_counter() - start
ratio = defective_time / patched_time if patched_time > 0 else float('inf')
print(f"\n N=100 init x 100 resp, {iters} iters:")
print(f" defective={defective_time:.4f}s patched={patched_time:.4f}s ratio={ratio:.1f}x")
self.assertGreater(
ratio, 3.0,
f"expected patched to be at least 3x faster at N=100 x 100, got {ratio:.1f}x"
)
if __name__ == "__main__":
unittest.main()