Skip to content

Commit 9f82f85

Browse files
committed
tests: avoid test port conflicts
Signed-off-by: Ryan Leung <rleungx@gmail.com>
1 parent 5c811b2 commit 9f82f85

18 files changed

Lines changed: 298 additions & 87 deletions

File tree

pkg/mcs/resourcemanager/server/server.go

Lines changed: 30 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -133,6 +133,19 @@ func (s *Server) Run() (err error) {
133133
if err = utils.InitClient(s); err != nil {
134134
return err
135135
}
136+
if err = s.initListenerAndUpdateConfig(); err != nil {
137+
return err
138+
}
139+
defer func() {
140+
if err != nil && s.serviceRegister != nil {
141+
if deregisterErr := s.serviceRegister.Deregister(); deregisterErr != nil {
142+
log.Warn("failed to deregister the service", errs.ZapError(deregisterErr))
143+
}
144+
}
145+
if err != nil && s.GetListener() != nil {
146+
_ = s.GetListener().Close()
147+
}
148+
}()
136149

137150
if s.serviceID, s.serviceRegister, err = utils.Register(s, constant.ResourceManagerServiceName); err != nil {
138151
return err
@@ -386,6 +399,23 @@ func (s *Server) GetTLSConfig() *grpcutil.TLSConfig {
386399
return &s.cfg.Security.TLSConfig
387400
}
388401

402+
func (s *Server) initListenerAndUpdateConfig() error {
403+
oldListenAddr := s.cfg.ListenAddr
404+
if err := s.InitListener(s.GetTLSConfig(), s.cfg.GetListenAddr()); err != nil {
405+
return err
406+
}
407+
actualListenAddr := s.GetActualListenAddr()
408+
if actualListenAddr == "" {
409+
return nil
410+
}
411+
s.cfg.ListenAddr = server.ResolveListenAddr(s.cfg.ListenAddr, actualListenAddr)
412+
s.cfg.AdvertiseListenAddr = server.ResolveAdvertiseListenAddr(s.cfg.AdvertiseListenAddr, actualListenAddr)
413+
if s.cfg.Name == "" || s.cfg.Name == oldListenAddr {
414+
s.cfg.Name = s.cfg.AdvertiseListenAddr
415+
}
416+
return nil
417+
}
418+
389419
// GetServingUrls gets service endpoints from the leader in election group.
390420
func (s *Server) GetServingUrls() []string {
391421
return s.participant.GetServingUrls()
@@ -416,9 +446,6 @@ func (s *Server) startServer() (err error) {
416446
rejectMetadataWritesViaGRPC: s.ShouldRejectMetadataWritesViaGRPC(),
417447
}
418448

419-
if err := s.InitListener(s.GetTLSConfig(), s.cfg.GetListenAddr()); err != nil {
420-
return err
421-
}
422449
// Only start the metering writer if a valid metering config is provided.
423450
if len(s.cfg.Metering.Type) > 0 {
424451
s.meteringWriter, err = metering.NewWriter(s.Context(), &s.cfg.Metering, fmt.Sprintf("pd%d", s.participant.ID()))

pkg/mcs/router/server/server.go

Lines changed: 30 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -120,6 +120,19 @@ func (s *Server) Run() (err error) {
120120
if err = utils.InitClient(s); err != nil {
121121
return err
122122
}
123+
if err = s.initListenerAndUpdateConfig(); err != nil {
124+
return err
125+
}
126+
defer func() {
127+
if err != nil && s.serviceRegister != nil {
128+
if deregisterErr := s.serviceRegister.Deregister(); deregisterErr != nil {
129+
log.Warn("failed to deregister the service", errs.ZapError(deregisterErr))
130+
}
131+
}
132+
if err != nil && s.GetListener() != nil {
133+
_ = s.GetListener().Close()
134+
}
135+
}()
123136

124137
if s.serviceID, s.serviceRegister, err = utils.Register(s, constant.RouterServiceName); err != nil {
125138
return err
@@ -181,6 +194,23 @@ func (s *Server) GetTLSConfig() *grpcutil.TLSConfig {
181194
return &s.cfg.Security.TLSConfig
182195
}
183196

197+
func (s *Server) initListenerAndUpdateConfig() error {
198+
oldListenAddr := s.cfg.ListenAddr
199+
if err := s.InitListener(s.GetTLSConfig(), s.cfg.GetListenAddr()); err != nil {
200+
return err
201+
}
202+
actualListenAddr := s.GetActualListenAddr()
203+
if actualListenAddr == "" {
204+
return nil
205+
}
206+
s.cfg.ListenAddr = server.ResolveListenAddr(s.cfg.ListenAddr, actualListenAddr)
207+
s.cfg.AdvertiseListenAddr = server.ResolveAdvertiseListenAddr(s.cfg.AdvertiseListenAddr, actualListenAddr)
208+
if s.cfg.Name == "" || s.cfg.Name == oldListenAddr {
209+
s.cfg.Name = s.cfg.AdvertiseListenAddr
210+
}
211+
return nil
212+
}
213+
184214
// GetCluster returns the cluster.
185215
func (s *Server) GetCluster() *Cluster {
186216
return s.cluster
@@ -235,9 +265,6 @@ func (s *Server) startServer() (err error) {
235265
bs.ServerMaxProcsGauge.Set(float64(runtime.GOMAXPROCS(0)))
236266

237267
s.service = &Service{Server: s}
238-
if err := s.InitListener(s.GetTLSConfig(), s.cfg.GetListenAddr()); err != nil {
239-
return err
240-
}
241268

242269
serverReadyChan := make(chan struct{})
243270
defer close(serverReadyChan)

pkg/mcs/scheduling/server/server.go

Lines changed: 30 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,19 @@ func (s *Server) Run() (err error) {
154154
if err = utils.InitClient(s); err != nil {
155155
return err
156156
}
157+
if err = s.initListenerAndUpdateConfig(); err != nil {
158+
return err
159+
}
160+
defer func() {
161+
if err != nil && s.serviceRegister != nil {
162+
if deregisterErr := s.serviceRegister.Deregister(); deregisterErr != nil {
163+
log.Warn("failed to deregister the service", errs.ZapError(deregisterErr))
164+
}
165+
}
166+
if err != nil && s.GetListener() != nil {
167+
_ = s.GetListener().Close()
168+
}
169+
}()
157170

158171
if s.serviceID, s.serviceRegister, err = utils.Register(s, constant.SchedulingServiceName); err != nil {
159172
return err
@@ -406,6 +419,23 @@ func (s *Server) GetTLSConfig() *grpcutil.TLSConfig {
406419
return &s.cfg.Security.TLSConfig
407420
}
408421

422+
func (s *Server) initListenerAndUpdateConfig() error {
423+
oldListenAddr := s.cfg.ListenAddr
424+
if err := s.InitListener(s.GetTLSConfig(), s.cfg.GetListenAddr()); err != nil {
425+
return err
426+
}
427+
actualListenAddr := s.GetActualListenAddr()
428+
if actualListenAddr == "" {
429+
return nil
430+
}
431+
s.cfg.ListenAddr = server.ResolveListenAddr(s.cfg.ListenAddr, actualListenAddr)
432+
s.cfg.AdvertiseListenAddr = server.ResolveAdvertiseListenAddr(s.cfg.AdvertiseListenAddr, actualListenAddr)
433+
if s.cfg.Name == "" || s.cfg.Name == oldListenAddr {
434+
s.cfg.Name = s.cfg.AdvertiseListenAddr
435+
}
436+
return nil
437+
}
438+
409439
// GetCluster returns the cluster.
410440
func (s *Server) GetCluster() *Cluster {
411441
cluster := s.cluster.Load()
@@ -478,9 +508,6 @@ func (s *Server) startServer() (err error) {
478508
s.service = &Service{Server: s}
479509
s.AddServiceReadyCallback(s.startCluster)
480510
s.AddServiceExitCallback(s.stopCluster)
481-
if err := s.InitListener(s.GetTLSConfig(), s.cfg.GetListenAddr()); err != nil {
482-
return err
483-
}
484511

485512
serverReadyChan := make(chan struct{})
486513
defer close(serverReadyChan)

pkg/mcs/server/server.go

Lines changed: 56 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,7 @@ type BaseServer struct {
4545
clientConns sync.Map
4646
secure bool
4747
muxListener net.Listener
48+
listenAddr string
4849
// Callback functions for different stages
4950
// startCallbacks will be called after the server is started.
5051
startCallbacks []func()
@@ -158,14 +159,41 @@ func (bs *BaseServer) InitListener(tlsCfg *grpcutil.TLSConfig, listenAddr string
158159
} else {
159160
bs.muxListener, err = net.Listen(constant.TCPNetworkStr, listenURL.Host)
160161
}
161-
return err
162+
if err != nil {
163+
return err
164+
}
165+
bs.listenAddr = buildActualListenAddr(listenURL, bs.muxListener.Addr())
166+
return nil
162167
}
163168

164169
// GetListener returns the listener.
165170
func (bs *BaseServer) GetListener() net.Listener {
166171
return bs.muxListener
167172
}
168173

174+
// GetActualListenAddr returns the listener address after binding.
175+
func (bs *BaseServer) GetActualListenAddr() string {
176+
return bs.listenAddr
177+
}
178+
179+
// ResolveListenAddr returns actualListenAddr only when listenAddr points at a
180+
// kernel-selected port.
181+
func ResolveListenAddr(listenAddr, actualListenAddr string) string {
182+
if hasZeroPort(listenAddr) {
183+
return actualListenAddr
184+
}
185+
return listenAddr
186+
}
187+
188+
// ResolveAdvertiseListenAddr returns actualListenAddr when advertiseAddr was left
189+
// unspecified or still points at a kernel-selected port.
190+
func ResolveAdvertiseListenAddr(advertiseAddr, actualListenAddr string) string {
191+
if advertiseAddr == "" || hasZeroPort(advertiseAddr) {
192+
return actualListenAddr
193+
}
194+
return advertiseAddr
195+
}
196+
169197
// IsSecure checks if the server enable TLS.
170198
func (bs *BaseServer) IsSecure() bool {
171199
return bs.secure
@@ -186,3 +214,30 @@ func (bs *BaseServer) CloseClientConns() {
186214
return true
187215
})
188216
}
217+
218+
func buildActualListenAddr(listenURL *url.URL, addr net.Addr) string {
219+
host, _, err := net.SplitHostPort(listenURL.Host)
220+
if err != nil || host == "" {
221+
host, _, _ = net.SplitHostPort(addr.String())
222+
}
223+
if ip := net.ParseIP(host); ip != nil && ip.IsUnspecified() {
224+
host = "127.0.0.1"
225+
}
226+
_, port, err := net.SplitHostPort(addr.String())
227+
if err != nil {
228+
return listenURL.String()
229+
}
230+
actualURL := *listenURL
231+
actualURL.Host = net.JoinHostPort(host, port)
232+
return actualURL.String()
233+
}
234+
235+
func hasZeroPort(addr string) bool {
236+
parsed, err := url.Parse(addr)
237+
host := addr
238+
if err == nil && parsed.Host != "" {
239+
host = parsed.Host
240+
}
241+
_, port, err := net.SplitHostPort(host)
242+
return err == nil && port == "0"
243+
}

pkg/mcs/server/server_test.go

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,57 @@
1+
// Copyright 2026 TiKV Project Authors.
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
// See the License for the specific language governing permissions and
13+
// limitations under the License.
14+
15+
package server
16+
17+
import (
18+
"context"
19+
"net/url"
20+
"testing"
21+
22+
"github.com/stretchr/testify/require"
23+
"go.uber.org/goleak"
24+
25+
"github.com/tikv/pd/pkg/utils/grpcutil"
26+
"github.com/tikv/pd/pkg/utils/testutil"
27+
)
28+
29+
func TestMain(m *testing.M) {
30+
goleak.VerifyTestMain(m, testutil.LeakOptions...)
31+
}
32+
33+
func TestInitListenerWithKernelSelectedPort(t *testing.T) {
34+
re := require.New(t)
35+
svr := NewBaseServer(context.Background())
36+
re.NoError(svr.InitListener(&grpcutil.TLSConfig{}, "http://127.0.0.1:0"))
37+
defer svr.GetListener().Close()
38+
39+
actual := svr.GetActualListenAddr()
40+
u, err := url.Parse(actual)
41+
re.NoError(err)
42+
re.Equal("http", u.Scheme)
43+
re.Equal("127.0.0.1", u.Hostname())
44+
re.NotEmpty(u.Port())
45+
re.NotEqual("0", u.Port())
46+
}
47+
48+
func TestResolveAdvertiseListenAddr(t *testing.T) {
49+
actualAddr := "http://127.0.0.1:12345"
50+
require.Equal(t, actualAddr, ResolveListenAddr("http://127.0.0.1:0", actualAddr))
51+
require.Equal(t, "http://127.0.0.1:23456", ResolveListenAddr("http://127.0.0.1:23456", actualAddr))
52+
53+
require.Equal(t, actualAddr, ResolveAdvertiseListenAddr("", actualAddr))
54+
require.Equal(t, actualAddr, ResolveAdvertiseListenAddr("http://127.0.0.1:0", actualAddr))
55+
require.Equal(t, actualAddr, ResolveAdvertiseListenAddr("127.0.0.1:0", actualAddr))
56+
require.Equal(t, "http://127.0.0.1:23456", ResolveAdvertiseListenAddr("http://127.0.0.1:23456", actualAddr))
57+
}

0 commit comments

Comments
 (0)