diff --git a/Makefile b/Makefile index 037e311..475e122 100644 --- a/Makefile +++ b/Makefile @@ -17,6 +17,9 @@ dtest: dfmt: docker run -v `pwd`:/gopath/src/gor -t -i gor go fmt +dvet: + docker run -v `pwd`:/gopath/src/gor -t -i gor go vet + dbench: docker run -v `pwd`:/gopath/src/gor -t -i gor go test -v -run NOT_EXISTING -bench HTTP diff --git a/input_raw.go b/input_raw.go index 3a6a960..20dd4b0 100644 --- a/input_raw.go +++ b/input_raw.go @@ -1,7 +1,7 @@ package main import ( - raw "github.com/buger/gor/raw_socket_listener" + raw "gor/raw_socket_listener" "log" "net" "strings" diff --git a/input_raw_test.go b/input_raw_test.go index b1bcca6..d446d11 100644 --- a/input_raw_test.go +++ b/input_raw_test.go @@ -1,8 +1,13 @@ package main import ( + "bytes" "io" + "io/ioutil" + "log" "net/http" + "net/http/httputil" + "os/exec" "strings" "sync" "testing" @@ -37,3 +42,123 @@ func TestRAWInput(t *testing.T) { close(quit) } + +func TestInputRAW100Expect(t *testing.T) { + wg := new(sync.WaitGroup) + quit := make(chan int) + + file_content, _ := ioutil.ReadFile("README.md") + + // Origing and Replay server initialization + origin := startHTTP(func(req *http.Request) { + defer req.Body.Close() + ioutil.ReadAll(req.Body) + + wg.Done() + }) + + origin_address := strings.Replace(origin.Addr().String(), "[::]", "127.0.0.1", -1) + + input := NewRAWInput(origin_address) + + // We will use it to get content of raw HTTP request + test_output := NewTestOutput(func(data []byte) { + if strings.Contains(string(data), "Expect: 100-continue") { + t.Error("Should not contain 100-continue header") + } + wg.Done() + }) + + listener := startHTTP(func(req *http.Request) { + defer req.Body.Close() + body, _ := ioutil.ReadAll(req.Body) + + if !bytes.Equal(body, file_content) { + buf, _ := httputil.DumpRequest(req, true) + t.Error("Wrong POST body:", string(buf)) + } + + wg.Done() + }) + replay_address := listener.Addr().String() + + headers := HTTPHeaders{HTTPHeader{"User-Agent", "Gor"}} + methods := HTTPMethods{"GET", "PUT", "POST"} + http_output := NewHTTPOutput(replay_address, headers, methods, HTTPUrlRegexp{}, HTTPHeaderFilters{}, HTTPHeaderHashFilters{}, "", UrlRewriteMap{}, 0) + + Plugins.Inputs = []io.Reader{input} + Plugins.Outputs = []io.Writer{test_output, http_output} + + go Start(quit) + + wg.Add(3) + curl := exec.Command("curl", "http://"+origin_address, "--data-binary", "@README.md") + err := curl.Run() + if err != nil { + log.Fatal(err) + } + + wg.Wait() + close(quit) +} + +func TestInputRAWChunkedEncoding(t *testing.T) { + wg := new(sync.WaitGroup) + quit := make(chan int) + + file_content, _ := ioutil.ReadFile("README.md") + + // Origing and Replay server initialization + origin := startHTTP(func(req *http.Request) { + defer req.Body.Close() + ioutil.ReadAll(req.Body) + + wg.Done() + }) + + origin_address := strings.Replace(origin.Addr().String(), "[::]", "127.0.0.1", -1) + + input := NewRAWInput(origin_address) + + // We will use it to get content of raw HTTP request + test_output := NewTestOutput(func(data []byte) { + if strings.Contains(string(data), "Expect: 100-continue") { + t.Error("Should not contain 100-continue header") + } + wg.Done() + }) + + listener := startHTTP(func(req *http.Request) { + defer req.Body.Close() + body, _ := ioutil.ReadAll(req.Body) + + if !bytes.Equal(body, file_content) { + buf, _ := httputil.DumpRequest(req, true) + t.Error("Wrong POST body:", string(buf)) + } + + wg.Done() + }) + replay_address := listener.Addr().String() + + headers := HTTPHeaders{HTTPHeader{"User-Agent", "Gor"}} + methods := HTTPMethods{"GET", "PUT", "POST"} + http_output := NewHTTPOutput(replay_address, headers, methods, HTTPUrlRegexp{}, HTTPHeaderFilters{}, HTTPHeaderHashFilters{}, "", UrlRewriteMap{}, 0) + + Plugins.Inputs = []io.Reader{input} + Plugins.Outputs = []io.Writer{test_output, http_output} + + go Start(quit) + + wg.Add(3) + + curl := exec.Command("curl", "http://"+origin_address, "--header", "Transfer-Encoding: chunked", "--data-binary", "@README.md") + err := curl.Run() + if err != nil { + log.Fatal(err) + } + + wg.Wait() + + close(quit) +} diff --git a/output_http.go b/output_http.go index 46a342a..2623641 100644 --- a/output_http.go +++ b/output_http.go @@ -7,7 +7,6 @@ import ( "io/ioutil" "log" "net/http" - "net/http/httputil" "net/url" "strings" "sync/atomic" @@ -32,9 +31,6 @@ func (o *HTTPOutput) customCheckRedirect(req *http.Request, via []*http.Request) func ParseRequest(data []byte) (request *http.Request, err error) { var body []byte - // Test if request have Transfer-Encoding: chunked - isChunked := bytes.Contains(data, []byte(": chunked\r\n")) - buf := bytes.NewBuffer(data) reader := bufio.NewReader(buf) @@ -46,12 +42,7 @@ func ParseRequest(data []byte) (request *http.Request, err error) { } if request.Method == "POST" { - // This works, because ReadRequest method modify buffer and strips all headers, leaving only body - if isChunked { - body, _ = ioutil.ReadAll(httputil.NewChunkedReader(reader)) - } else { - body, _ = ioutil.ReadAll(reader) - } + body, _ = ioutil.ReadAll(reader) bodyBuf := bytes.NewBuffer(body) diff --git a/output_http_test.go b/output_http_test.go index 2432d5c..c954cc3 100644 --- a/output_http_test.go +++ b/output_http_test.go @@ -6,7 +6,6 @@ import ( "net" "net/http" "net/http/httputil" - _ "strings" "sync" "testing" "time" @@ -51,9 +50,6 @@ func TestHTTPOutput(t *testing.T) { input := NewTestInput() - headers := HTTPHeaders{HTTPHeader{"User-Agent", "Gor"}} - methods := HTTPMethods{"GET", "PUT", "POST"} - listener := startHTTP(func(req *http.Request) { if req.Header.Get("User-Agent") != "Gor" { t.Error("Wrong header") @@ -76,6 +72,9 @@ func TestHTTPOutput(t *testing.T) { wg.Done() }) + headers := HTTPHeaders{HTTPHeader{"User-Agent", "Gor"}} + methods := HTTPMethods{"GET", "PUT", "POST"} + output := NewHTTPOutput(listener.Addr().String(), headers, methods, HTTPUrlRegexp{}, HTTPHeaderFilters{}, HTTPHeaderHashFilters{}, "", UrlRewriteMap{}, 0) Plugins.Inputs = []io.Reader{input} @@ -95,42 +94,6 @@ func TestHTTPOutput(t *testing.T) { close(quit) } -func TestHTTPOutputChunkedEncoding(t *testing.T) { - wg := new(sync.WaitGroup) - quit := make(chan int) - - input := NewTestInput() - - headers := HTTPHeaders{HTTPHeader{"User-Agent", "Gor"}} - methods := HTTPMethods{"GET", "PUT", "POST"} - - listener := startHTTP(func(req *http.Request) { - defer req.Body.Close() - body, _ := ioutil.ReadAll(req.Body) - - if string(body) != "Wikipedia in\r\n\r\nchunks." { - buf, _ := httputil.DumpRequest(req, true) - t.Error("Wrong POST body:", buf, body, []byte("Wikipedia in\r\n\r\nchunks.")) - } - - wg.Done() - }) - - output := NewHTTPOutput(listener.Addr().String(), headers, methods, HTTPUrlRegexp{}, HTTPHeaderFilters{}, HTTPHeaderHashFilters{}, "", UrlRewriteMap{}, 0) - - Plugins.Inputs = []io.Reader{input} - Plugins.Outputs = []io.Writer{output} - - go Start(quit) - - wg.Add(1) - input.EmitChunkedPOST() - - wg.Wait() - - close(quit) -} - func BenchmarkHTTPOutput(b *testing.B) { wg := new(sync.WaitGroup) quit := make(chan int) diff --git a/output_tcp_test.go b/output_tcp_test.go index 045ccfe..aaf0cf6 100644 --- a/output_tcp_test.go +++ b/output_tcp_test.go @@ -44,6 +44,7 @@ func startTCP(cb func([]byte)) net.Listener { go func() { for { conn, _ := listener.Accept() + defer conn.Close() go func() { reader := bufio.NewReader(conn) @@ -59,7 +60,6 @@ func startTCP(cb func([]byte)) net.Listener { } cb(new_buf) } - conn.Close() }() } }() diff --git a/raw_socket_listener/listener.go b/raw_socket_listener/listener.go index 625a100..b5f7150 100644 --- a/raw_socket_listener/listener.go +++ b/raw_socket_listener/listener.go @@ -5,6 +5,7 @@ import ( "log" "net" "strconv" + "bytes" ) // Capture traffic from socket using RAW_SOCKET's @@ -17,6 +18,11 @@ import ( type Listener struct { messages map[string]*TCPMessage // buffer of TCPMessages waiting to be send + // Expect: 100-continue request is send in 2 tcp messages + // We store ACK aliases to merge this packets together + ack_aliases map[uint32]uint32 + seq_with_data map[uint32]uint32 + c_packets chan *TCPPacket c_messages chan *TCPMessage // Messages ready to be send to client @@ -30,10 +36,13 @@ type Listener struct { func NewListener(addr string, port string) (rawListener *Listener) { rawListener = &Listener{} - rawListener.c_packets = make(chan *TCPPacket, 100) - rawListener.c_messages = make(chan *TCPMessage, 100) - rawListener.c_del_message = make(chan *TCPMessage, 100) + rawListener.c_packets = make(chan *TCPPacket, 10000) + rawListener.c_messages = make(chan *TCPMessage, 10000) + rawListener.c_del_message = make(chan *TCPMessage, 10000) + rawListener.messages = make(map[string]*TCPMessage) + rawListener.ack_aliases = make(map[uint32]uint32) + rawListener.seq_with_data = make(map[uint32]uint32) rawListener.addr = addr rawListener.port, _ = strconv.Atoi(port) @@ -50,6 +59,7 @@ func (t *Listener) listen() { // If message ready for deletion it means that its also complete or expired by timeout case message := <-t.c_del_message: t.c_messages <- message + delete(t.ack_aliases, message.packets[0].Ack) delete(t.messages, message.ID) // We need to use channels to process each packet to avoid data races @@ -68,7 +78,7 @@ func (t *Listener) readRAWSocket() { defer conn.Close() - buf := make([]byte, 4096*2) + buf := make([]byte, 4096*10) for { // Note: ReadFrom receive messages without IP header @@ -115,6 +125,9 @@ func (t *Listener) isIncomingDataPacket(buf []byte) bool { return false } +var bExpect100ContinueCheck = []byte("Expect: 100-continue") +var bPOST = []byte("POST") + // Trying to add packet to existing message or creating new message // // For TCP message unique id is Acknowledgment number (see tcp_packet.go) @@ -122,8 +135,19 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { defer func() { recover() }() var message *TCPMessage - m_id := packet.Addr.String() + strconv.Itoa(int(packet.Ack)) + parent_message_ack, parent_ok := t.seq_with_data[packet.Seq] + if parent_ok { + t.ack_aliases[packet.Ack] = parent_message_ack + delete(t.seq_with_data, packet.Seq) + } + + ack_alias, alias_ok := t.ack_aliases[packet.Ack] + if alias_ok { + packet.Ack = ack_alias + } + + m_id := packet.Addr.String() + strconv.Itoa(int(packet.Ack)) message, ok := t.messages[m_id] if !ok { @@ -132,6 +156,15 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { t.messages[m_id] = message } + if bytes.Equal(packet.Data[0:4], bPOST) { + if bytes.Equal(packet.Data[len(packet.Data)-24:len(packet.Data)-4], bExpect100ContinueCheck) { + t.seq_with_data[packet.Seq + uint32(len(packet.Data))] = packet.Ack + + // Removing `Expect: 100-continue` header + packet.Data = append(packet.Data[:len(packet.Data)-24], packet.Data[len(packet.Data)-2:]...) + } + } + // Adding packet to message message.c_packets <- packet } diff --git a/raw_socket_listener/tcp_message.go b/raw_socket_listener/tcp_message.go index 21ed63e..9e45a85 100644 --- a/raw_socket_listener/tcp_message.go +++ b/raw_socket_listener/tcp_message.go @@ -4,6 +4,10 @@ import ( "log" "sort" "time" + "bytes" + "net/http/httputil" + "bufio" + "io/ioutil" ) const MSG_EXPIRE = 2000 * time.Millisecond @@ -71,6 +75,30 @@ func (t *TCPMessage) Timeout() { } } +var bTransferEncodingChunked = []byte("Transfer-Encoding: chunked\r\n") +var b2xCRLF = []byte("\r\n\r\n") + +// Norimalize requests with `Transfer-Encoding: chunked` header, because they have special body format +func fixChunkedEncoding(data []byte) []byte { + if bytes.Equal(data[0:4], bPOST) { + body_idx := bytes.Index(data, b2xCRLF) + chunked_header_idx := bytes.Index(data[:body_idx], bTransferEncodingChunked) + + if chunked_header_idx != -1 { + buf := bytes.NewBuffer(data[body_idx+4:]) + // Adding 4 bytes to skip 2xCLRF + bodyReader := bufio.NewReader(buf) + body, _ := ioutil.ReadAll(httputil.NewChunkedReader(bodyReader)) + + // Exclude Transfer-Encoding header and append new body + return append(append(append(data[:chunked_header_idx], + data[chunked_header_idx+len(bTransferEncodingChunked):body_idx]...), b2xCRLF...), body...) + } + } + + return data +} + // Bytes sorts packets in right orders and return message content func (t *TCPMessage) Bytes() (output []byte) { sort.Sort(BySeq(t.packets)) @@ -79,7 +107,7 @@ func (t *TCPMessage) Bytes() (output []byte) { output = append(output, v.Data...) } - return + return fixChunkedEncoding(output) } // AddPacket to the message and ensure packet uniqueness diff --git a/raw_socket_listener/tcp_packet.go b/raw_socket_listener/tcp_packet.go index 5e81ff5..5af8a73 100644 --- a/raw_socket_listener/tcp_packet.go +++ b/raw_socket_listener/tcp_packet.go @@ -89,6 +89,7 @@ func (t *TCPPacket) String() string { "Window size:" + strconv.Itoa(int(t.Window)), "Checksum:" + strconv.Itoa(int(t.Checksum)), + "Data size:" + strconv.Itoa(len(t.Data)), "Data:" + string(t.Data), }, "\n") }