Merge pull request #158 from buger/fix-100-expect

Fix Expect: 100-continue requests and refactor chunked encoding
This commit is contained in:
Leonid Bugaev
2015-06-29 17:29:59 +05:00
8 changed files with 201 additions and 57 deletions
+3
View File
@@ -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
+125
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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()
}()
}
}()
+38 -5
View File
@@ -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
}
+29 -1
View File
@@ -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
+1
View File
@@ -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")
}