Skip to content

Commit 906cf7c

Browse files
committed
[testing/freeport] fix race condition in port reservation
- change resolverFn to return release function instead of using defer - call release() only after checking/updating the reserved map - rename getPort to findFreePort, tcpResolv to tcpResolver - remove redundant comments Change-Id: I26b6dab633f95a788e53f5e08b1a1e2b4697b82b
1 parent 906b77a commit 906cf7c

2 files changed

Lines changed: 49 additions & 37 deletions

File tree

testing/freeport/freeport.go

Lines changed: 49 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -12,28 +12,35 @@ import (
1212
// It maintains a registry of ports that are currently reserved by tests
1313
// and provides methods to get free ports or reserve ports exclusively.
1414
type resolver struct {
15-
// reserved maps port numbers to the test that reserved them
16-
reserved map[int]*testing.T
17-
// reservedMu protects access to the reserved map
15+
reserved map[int]*testing.T
1816
reservedMu sync.Mutex `exhaustruct:"optional"`
1917

20-
// resolverFn is the function used to find free ports for this protocol
21-
resolverFn func() (int, error)
18+
// resolverFn finds a free port and returns it along with a release
19+
// function. The release function MUST be called to free the underlying
20+
// resource (listener/connection) that holds the port. This design
21+
// prevents race conditions where another probe could grab the same
22+
// port between when it's discovered and when it's added to the map.
23+
resolverFn func() (port int, release func(), err error)
2224
}
2325

24-
func (r *resolver) getPort() (int, error) {
26+
func (r *resolver) findFreePort() (int, error) {
2527
r.reservedMu.Lock()
2628
defer r.reservedMu.Unlock()
2729

2830
for {
29-
port, err := r.resolverFn()
31+
port, release, err := r.resolverFn()
3032
if err != nil {
3133
return 0, e.NewFrom("failed to get free port", err)
3234
}
3335

34-
if _, ok := r.reserved[port]; !ok {
35-
return port, nil
36+
if _, reserved := r.reserved[port]; reserved {
37+
release()
38+
continue
3639
}
40+
41+
release()
42+
43+
return port, nil
3744
}
3845
}
3946

@@ -48,19 +55,22 @@ func (r *resolver) reservePort(t *testing.T) (int, error) { //nolint:thelper
4855
defer r.reservedMu.Unlock()
4956

5057
for {
51-
port, err := r.resolverFn()
58+
port, release, err := r.resolverFn()
5259
if err != nil {
5360
return 0, e.NewFrom("failed to reserve free port", err)
5461
}
5562

56-
if _, ok := r.reserved[port]; !ok {
57-
r.reserved[port] = t
58-
t.Cleanup(func() {
59-
r.releasePort(t, port)
60-
})
61-
62-
return port, nil
63+
if _, reserved := r.reserved[port]; reserved {
64+
release()
65+
continue
6366
}
67+
68+
r.reserved[port] = t
69+
t.Cleanup(func() { r.releasePort(t, port) })
70+
71+
release()
72+
73+
return port, nil
6474
}
6575
}
6676

@@ -83,38 +93,42 @@ func (r *resolver) releasePort(t *testing.T, port int) bool { //nolint:thelper
8393
}
8494

8595
var (
86-
tcpResolv = &resolver{ //nolint:gochecknoglobals
96+
tcpResolver = &resolver{ //nolint:gochecknoglobals
8797
reserved: make(map[int]*testing.T),
88-
resolverFn: func() (int, error) {
98+
resolverFn: func() (int, func(), error) {
8999
addr, err := net.ResolveTCPAddr("tcp", "localhost:0")
90100
if err != nil {
91-
return 0, e.NewFrom("failed to resolve TCP address", err)
101+
return 0, nil, e.NewFrom("failed to resolve TCP address", err)
92102
}
93103

94104
listener, err := net.ListenTCP("tcp", addr)
95105
if err != nil {
96-
return 0, e.NewFrom("failed to listen on TCP port", err)
106+
return 0, nil, e.NewFrom("failed to listen on TCP port", err)
97107
}
98-
defer func() { _ = listener.Close() }()
99108

100-
return listener.Addr().(*net.TCPAddr).Port, nil //nolint:forcetypeassert
109+
port := listener.Addr().(*net.TCPAddr).Port //nolint:forcetypeassert
110+
release := func() { _ = listener.Close() }
111+
112+
return port, release, nil
101113
},
102114
}
103-
udpResolv = &resolver{ //nolint:gochecknoglobals
115+
udpResolver = &resolver{ //nolint:gochecknoglobals
104116
reserved: make(map[int]*testing.T),
105-
resolverFn: func() (int, error) {
117+
resolverFn: func() (int, func(), error) {
106118
addr, err := net.ResolveUDPAddr("udp", "localhost:0")
107119
if err != nil {
108-
return 0, e.NewFrom("failed to resolve UDP address", err)
120+
return 0, nil, e.NewFrom("failed to resolve UDP address", err)
109121
}
110122

111123
conn, err := net.ListenUDP("udp", addr)
112124
if err != nil {
113-
return 0, e.NewFrom("failed to listen on UDP port", err)
125+
return 0, nil, e.NewFrom("failed to listen on UDP port", err)
114126
}
115-
defer func() { _ = conn.Close() }()
116127

117-
return conn.LocalAddr().(*net.UDPAddr).Port, nil //nolint:forcetypeassert
128+
port := conn.LocalAddr().(*net.UDPAddr).Port //nolint:forcetypeassert
129+
release := func() { _ = conn.Close() }
130+
131+
return port, release, nil
118132
},
119133
}
120134
)
@@ -123,14 +137,14 @@ var (
123137
// This function does not reserve the port, so there's a small chance
124138
// another process could claim it before you use it.
125139
func TCP() (int, error) {
126-
return tcpResolv.getPort()
140+
return tcpResolver.findFreePort()
127141
}
128142

129143
// UDP returns a free UDP port that is available for use.
130144
// This function does not reserve the port, so there's a small chance
131145
// another process could claim it before you use it.
132146
func UDP() (int, error) {
133-
return udpResolv.getPort()
147+
return udpResolver.findFreePort()
134148
}
135149

136150
// ReserveTCP reserves a TCP port for exclusive use within the test.
@@ -141,7 +155,7 @@ func UDP() (int, error) {
141155
func ReserveTCP(t *testing.T) int {
142156
t.Helper()
143157

144-
port, err := tcpResolv.reservePort(t)
158+
port, err := tcpResolver.reservePort(t)
145159
if err != nil {
146160
t.Fatalf("ReserveTCP: %v", err)
147161
}
@@ -157,7 +171,7 @@ func ReserveTCP(t *testing.T) int {
157171
func ReserveUDP(t *testing.T) int {
158172
t.Helper()
159173

160-
port, err := udpResolv.reservePort(t)
174+
port, err := udpResolver.reservePort(t)
161175
if err != nil {
162176
t.Fatalf("ReserveUDP: %v", err)
163177
}
@@ -174,7 +188,7 @@ func ReserveUDP(t *testing.T) int {
174188
// Returns true if the port was released, false otherwise.
175189
func ReleaseTCP(t *testing.T, port int) bool {
176190
t.Helper()
177-
return tcpResolv.releasePort(t, port)
191+
return tcpResolver.releasePort(t, port)
178192
}
179193

180194
// ReleaseUDP manually releases a UDP port that was reserved for the test. This
@@ -186,5 +200,5 @@ func ReleaseTCP(t *testing.T, port int) bool {
186200
// Returns true if the port was released, false otherwise.
187201
func ReleaseUDP(t *testing.T, port int) bool {
188202
t.Helper()
189-
return udpResolv.releasePort(t, port)
203+
return udpResolver.releasePort(t, port)
190204
}

testing/freeport/freeport_test.go

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,6 @@ func TestGet(t *testing.T) {
2020
require.NoError(t, err)
2121
require.Positive(t, p)
2222

23-
// Verify we can listen on the port
2423
l, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1"), Port: p, Zone: ""})
2524
require.NoError(t, err)
2625
require.NotNil(t, l)
@@ -36,7 +35,6 @@ func TestGet(t *testing.T) {
3635
require.NoError(t, err)
3736
require.Positive(t, p)
3837

39-
// Verify we can bind UDP to the port
4038
c, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: p, Zone: ""})
4139
require.NoError(t, err)
4240
require.NotNil(t, c)

0 commit comments

Comments
 (0)