Skip to content

Commit 153a5ba

Browse files
committed
Fix websocket PR lint issues and add coverage
1 parent 60efa9a commit 153a5ba

4 files changed

Lines changed: 36 additions & 10 deletions

File tree

dialer.go

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -33,11 +33,11 @@ type Dialer interface {
3333
Dial(ctx context.Context, session *session, attempt int, tlsConfig *tls.Config) (net.Conn, error)
3434
}
3535

36-
type TcpDialer struct {
36+
type TCPDialer struct {
3737
ctxDialer proxy.ContextDialer
3838
}
3939

40-
func (d *TcpDialer) Dial(ctx context.Context, session *session, attempt int, tlsConfig *tls.Config) (conn net.Conn, err error) {
40+
func (d *TCPDialer) Dial(ctx context.Context, session *session, attempt int, tlsConfig *tls.Config) (conn net.Conn, err error) {
4141
address := session.SocketConnectAddress[attempt%len(session.SocketConnectAddress)]
4242
session.log.OnEventf("Connecting to: %v", address)
4343

@@ -107,7 +107,7 @@ func loadDialerConfig(settings *SessionSettings) (dialer Dialer, err error) {
107107
}
108108

109109
stdDialer := &net.Dialer{}
110-
dialer = &TcpDialer{
110+
dialer = &TCPDialer{
111111
ctxDialer: stdDialer,
112112
}
113113
if settings.HasSetting(config.SocketTimeout) {
@@ -163,7 +163,7 @@ func loadDialerConfig(settings *SessionSettings) (dialer Dialer, err error) {
163163
}
164164

165165
if contextDialer, ok := proxyDialer.(proxy.ContextDialer); ok {
166-
dialer = &TcpDialer{
166+
dialer = &TCPDialer{
167167
ctxDialer: contextDialer,
168168
}
169169
} else {

dialer_test.go

Lines changed: 22 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,7 @@ func (s *DialerTestSuite) TestLoadDialerNoSettings() {
4242
dialer, err := loadDialerConfig(s.settings.GlobalSettings())
4343
s.Require().Nil(err)
4444

45-
stdDialer, ok := dialer.(*TcpDialer).ctxDialer.(*net.Dialer)
45+
stdDialer, ok := dialer.(*TCPDialer).ctxDialer.(*net.Dialer)
4646
s.Require().True(ok)
4747
s.Require().NotNil(stdDialer)
4848
s.Zero(stdDialer.Timeout)
@@ -53,7 +53,7 @@ func (s *DialerTestSuite) TestLoadDialerWithTimeout() {
5353
dialer, err := loadDialerConfig(s.settings.GlobalSettings())
5454
s.Require().Nil(err)
5555

56-
stdDialer, ok := dialer.(*TcpDialer).ctxDialer.(*net.Dialer)
56+
stdDialer, ok := dialer.(*TCPDialer).ctxDialer.(*net.Dialer)
5757
s.Require().True(ok)
5858
s.Require().NotNil(stdDialer)
5959
s.EqualValues(10*time.Second, stdDialer.Timeout)
@@ -73,7 +73,7 @@ func (s *DialerTestSuite) TestLoadDialerSocksProxy() {
7373
s.Require().Nil(err)
7474
s.Require().NotNil(dialer)
7575

76-
_, ok := dialer.(*TcpDialer).ctxDialer.(*net.Dialer)
76+
_, ok := dialer.(*TCPDialer).ctxDialer.(*net.Dialer)
7777
s.Require().False(ok)
7878
}
7979

@@ -90,3 +90,22 @@ func (s *DialerTestSuite) TestLoadDialerSocksProxyInvalidPort() {
9090
_, err := loadDialerConfig(s.settings.GlobalSettings())
9191
s.Require().NotNil(err)
9292
}
93+
94+
func (s *DialerTestSuite) TestLoadDialerWebsocket() {
95+
s.settings.GlobalSettings().Set(config.WebsocketLocation, "ws://example.com/ws")
96+
s.settings.GlobalSettings().Set(config.WebsocketOrigin, "http://localhost/")
97+
98+
dialer, err := loadDialerConfig(s.settings.GlobalSettings())
99+
s.Require().NoError(err)
100+
101+
wsDialer, ok := dialer.(*WebsocketDialer)
102+
s.Require().True(ok)
103+
s.Equal("ws://example.com/ws", wsDialer.wsConfig.Location.String())
104+
s.Equal("http://localhost/", wsDialer.wsConfig.Origin.String())
105+
}
106+
107+
func (s *DialerTestSuite) TestLoadDialerWebsocketMissingOrigin() {
108+
s.settings.GlobalSettings().Set(config.WebsocketLocation, "ws://example.com/ws")
109+
_, err := loadDialerConfig(s.settings.GlobalSettings())
110+
s.Require().Error(err)
111+
}

initiator.go

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -181,10 +181,8 @@ func (i *Initiator) handleConnection(session *session, tlsConfig *tls.Config, di
181181
if err != nil {
182182
session.log.OnEventf("Failed to connect: %v", err)
183183
goto reconnect
184-
} else {
185-
address := netConn.RemoteAddr().String()
186-
session.log.OnEventf("connected to remote address: %v", address)
187184
}
185+
session.log.OnEventf("connected to remote address: %v", netConn.RemoteAddr().String())
188186

189187
msgIn = make(chan fixIn, session.InChanCapacity)
190188
msgOut = make(chan []byte)

session_factory_test.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -649,6 +649,15 @@ func (s *SessionFactorySuite) TestConfigureSocketConnectAddress() {
649649
}
650650
}
651651

652+
func (s *SessionFactorySuite) TestConfigureSocketConnectAddressWebsocketOnly() {
653+
sess := new(session)
654+
s.SessionSettings.Set(config.WebsocketLocation, "wss://example.com/ws")
655+
656+
err := s.configureSocketConnectAddress(sess, s.SessionSettings)
657+
s.Require().NoError(err)
658+
s.Empty(sess.SocketConnectAddress)
659+
}
660+
652661
func (s *SessionFactorySuite) TestConfigureSocketConnectAddressMulti() {
653662
session := new(session)
654663
s.SessionSettings.Set(config.SocketConnectHost, "127.0.0.1")

0 commit comments

Comments
 (0)