fix(httphandler): post/put request hang (#154)
SonarQube Scan / SonarQube Trigger (push) Successful in 3m57s
SonarQube Scan / SonarQube Trigger (push) Successful in 3m57s
#153 Reviewed-on: #154 Co-authored-by: Bagas <bagas@fossy.my.id> Co-committed-by: Bagas <bagas@fossy.my.id>
This commit was merged in pull request #154.
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
package transport
|
package transport
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -52,25 +53,43 @@ func (hh *httpHandler) badRequest(conn net.Conn) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func readHTTPHeader(br *bufio.Reader, limit int) ([]byte, error) {
|
||||||
|
var headerBuf []byte
|
||||||
|
for {
|
||||||
|
line, err := br.ReadSlice('\n')
|
||||||
|
headerBuf = append(headerBuf, line...)
|
||||||
|
if errors.Is(err, bufio.ErrBufferFull) {
|
||||||
|
if len(headerBuf) > limit {
|
||||||
|
return nil, fmt.Errorf("headers too large")
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if bytes.HasSuffix(headerBuf, []byte("\r\n\r\n")) {
|
||||||
|
return headerBuf, nil
|
||||||
|
}
|
||||||
|
if len(headerBuf) > limit {
|
||||||
|
return nil, fmt.Errorf("headers too large")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
|
func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
|
||||||
defer hh.closeConnection(conn)
|
defer hh.closeConnection(conn)
|
||||||
|
|
||||||
_ = conn.SetReadDeadline(time.Now().Add(10 * time.Second))
|
_ = conn.SetReadDeadline(time.Now().Add(10 * time.Second))
|
||||||
buf := make([]byte, hh.config.HeaderSize())
|
br := bufio.NewReaderSize(conn, hh.config.HeaderSize())
|
||||||
n, err := conn.Read(buf)
|
headerBuf, err := readHTTPHeader(br, hh.config.HeaderSize())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = hh.badRequest(conn)
|
_ = hh.badRequest(conn)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if idx := bytes.Index(buf[:n], []byte("\r\n\r\n")); idx == -1 {
|
|
||||||
_ = hh.badRequest(conn)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
_ = conn.SetReadDeadline(time.Time{})
|
_ = conn.SetReadDeadline(time.Time{})
|
||||||
|
|
||||||
reqhf, err := header.NewRequest(buf[:n])
|
reqhf, err := header.NewRequest(headerBuf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Printf("Error creating request header: %v", err)
|
log.Printf("Error creating request header: %v", err)
|
||||||
_ = hh.badRequest(conn)
|
_ = hh.badRequest(conn)
|
||||||
@@ -101,7 +120,7 @@ func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
hw := stream.New(conn, conn, conn.RemoteAddr())
|
hw := stream.New(conn, br, conn.RemoteAddr())
|
||||||
defer func(hw stream.HTTP) {
|
defer func(hw stream.HTTP) {
|
||||||
err = hw.Close()
|
err = hw.Close()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -321,8 +321,14 @@ func TestHandler(t *testing.T) {
|
|||||||
isTLS: false,
|
isTLS: false,
|
||||||
redirectTLS: false,
|
redirectTLS: false,
|
||||||
request: []byte(""),
|
request: []byte(""),
|
||||||
expected: []byte("HTTP/1.1 400 Bad Request\r\n\r\n"),
|
expected: []byte(""),
|
||||||
setupMocks: func(msr *MockSessionRegistry) {
|
setupConn: func() (net.Conn, net.Conn) {
|
||||||
|
mc := new(MockConn)
|
||||||
|
mc.ReadBuffer = bytes.NewBuffer(nil)
|
||||||
|
mc.On("SetReadDeadline", mock.Anything).Return(nil)
|
||||||
|
mc.On("Write", []byte("HTTP/1.1 400 Bad Request\r\n\r\n")).Return(0, nil)
|
||||||
|
mc.On("Close").Return(nil)
|
||||||
|
return mc, nil
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -715,3 +721,113 @@ func TestHandler(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHandlerForwardsPostBody(t *testing.T) {
|
||||||
|
mockSessionRegistry := new(MockSessionRegistry)
|
||||||
|
mockConfig := &MockConfig{}
|
||||||
|
mockConfig.On("Domain").Return("example.com")
|
||||||
|
mockConfig.On("HTTPPort").Return("0")
|
||||||
|
mockConfig.On("HeaderSize").Return(4096)
|
||||||
|
mockConfig.On("TLSRedirect").Return(true)
|
||||||
|
hh := &httpHandler{
|
||||||
|
sessionRegistry: mockSessionRegistry,
|
||||||
|
config: mockConfig,
|
||||||
|
}
|
||||||
|
|
||||||
|
mockSession := new(MockSession)
|
||||||
|
mockForwarder := new(MockForwarder)
|
||||||
|
mockSSHChannel := new(MockSSHChannel)
|
||||||
|
|
||||||
|
mockSessionRegistry.On("Get", types.SessionKey{
|
||||||
|
Id: "test",
|
||||||
|
Type: types.TunnelTypeHTTP,
|
||||||
|
}).Return(mockSession, nil)
|
||||||
|
mockSession.On("Forwarder").Return(mockForwarder)
|
||||||
|
|
||||||
|
reqCh := make(chan *ssh.Request)
|
||||||
|
mockForwarder.On("OpenForwardedChannel", mock.Anything, mock.Anything).Return(mockSSHChannel, (<-chan *ssh.Request)(reqCh), nil)
|
||||||
|
|
||||||
|
var mu sync.Mutex
|
||||||
|
var capturedHeaders []byte
|
||||||
|
mockSSHChannel.On("Write", mock.Anything).Run(func(args mock.Arguments) {
|
||||||
|
mu.Lock()
|
||||||
|
capturedHeaders = append(capturedHeaders, args.Get(0).([]byte)...)
|
||||||
|
mu.Unlock()
|
||||||
|
}).Return(0, nil)
|
||||||
|
mockSSHChannel.On("Close").Return(nil)
|
||||||
|
|
||||||
|
bodyChan := make(chan string, 1)
|
||||||
|
mockForwarder.On("HandleConnection", mock.Anything, mockSSHChannel).Run(func(args mock.Arguments) {
|
||||||
|
w := args.Get(0).(io.ReadWriter)
|
||||||
|
buf := make([]byte, len("hello=world"))
|
||||||
|
if _, err := io.ReadFull(w, buf); err != nil {
|
||||||
|
bodyChan <- ""
|
||||||
|
} else {
|
||||||
|
bodyChan <- string(buf)
|
||||||
|
}
|
||||||
|
_, _ = w.Write([]byte("HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok"))
|
||||||
|
})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
for range reqCh {
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
serverConn, clientConn := net.Pipe()
|
||||||
|
defer func() {
|
||||||
|
_ = clientConn.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
remoteAddr, _ := net.ResolveTCPAddr("tcp", "127.0.0.1:12345")
|
||||||
|
wrappedServerConn := &wrappedConn{Conn: serverConn, remoteAddr: remoteAddr}
|
||||||
|
|
||||||
|
go hh.Handler(wrappedServerConn, true)
|
||||||
|
|
||||||
|
request := []byte("POST / HTTP/1.1\r\nHost: test.domain\r\nContent-Type: application/x-www-form-urlencoded\r\nContent-Length: 11\r\n\r\nhello=world")
|
||||||
|
go func() {
|
||||||
|
_, _ = clientConn.Write(request)
|
||||||
|
}()
|
||||||
|
|
||||||
|
var response []byte
|
||||||
|
respDone := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(respDone)
|
||||||
|
buf := make([]byte, 4096)
|
||||||
|
for {
|
||||||
|
n, err := clientConn.Read(buf)
|
||||||
|
if err != nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
response = append(response, buf[:n]...)
|
||||||
|
if bytes.Contains(response, []byte("\r\n\r\nok")) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case body := <-bodyChan:
|
||||||
|
assert.Equal(t, "hello=world", body)
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for forwarded body")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-respDone:
|
||||||
|
resStr := string(response)
|
||||||
|
assert.True(t, strings.HasPrefix(resStr, "HTTP/1.1 200 OK\r\n"))
|
||||||
|
assert.Contains(t, resStr, "Server: Tunnel Please\r\n")
|
||||||
|
assert.True(t, strings.HasSuffix(resStr, "\r\n\r\nok"))
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for response")
|
||||||
|
}
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
hdrStr := string(capturedHeaders)
|
||||||
|
mu.Unlock()
|
||||||
|
assert.Contains(t, hdrStr, "POST / HTTP/1.1\r\n")
|
||||||
|
assert.Contains(t, hdrStr, "Content-Length: 11\r\n")
|
||||||
|
assert.Contains(t, hdrStr, "X-Forwarded-For: 127.0.0.1\r\n")
|
||||||
|
|
||||||
|
mockSessionRegistry.AssertExpectations(t)
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user