mirror of
https://github.com/buger/goreplay.git
synced 2024-04-21 12:32:02 +00:00
Fix 100-Expect requests and refactor chunked encoding
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
raw "github.com/buger/gor/raw_socket_listener"
|
||||
raw "gor/raw_socket_listener"
|
||||
"log"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+1
-10
@@ -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)
|
||||
|
||||
|
||||
+3
-40
@@ -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)
|
||||
|
||||
+1
-1
@@ -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()
|
||||
}()
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user