diff --git a/.github/workflows/ci-test.yaml b/.github/workflows/ci-test.yaml index 80f9820..a099b08 100644 --- a/.github/workflows/ci-test.yaml +++ b/.github/workflows/ci-test.yaml @@ -26,4 +26,4 @@ jobs: restore-keys: | ${{ runner.os }}-go- - name: test - run: sudo go test ./... -v -timeout 120s -race + run: sudo go test ./... -v -timeout 120s diff --git a/examples/middleware/echo.sh b/examples/middleware/echo.sh index e9cb46b..ad22956 100755 --- a/examples/middleware/echo.sh +++ b/examples/middleware/echo.sh @@ -8,7 +8,7 @@ # function log { - if $GOR_TEST != ""; then # if we are not testing + if [[ ! -v GOR_TEST ]]; then # if we are not testing # Logging to stderr, because stdout/stdin used for data transfer >&2 echo "[DEBUG][ECHO] $1" fi diff --git a/examples/middleware/token_modifier.go b/examples/middleware/token_modifier.go index 01cb495..5367aba 100644 --- a/examples/middleware/token_modifier.go +++ b/examples/middleware/token_modifier.go @@ -115,7 +115,7 @@ func encode(buf []byte) []byte { } func Debug(args ...interface{}) { - if os.Getenv("GOR_TEST") != "" { // if we are not testing + if os.Getenv("GOR_TEST") == "" { // if we are not testing fmt.Fprint(os.Stderr, "[DEBUG][TOKEN-MOD] ") fmt.Fprintln(os.Stderr, args...) } diff --git a/middleware.go b/middleware.go index bd469af..5f775a4 100644 --- a/middleware.go +++ b/middleware.go @@ -6,10 +6,11 @@ import ( "encoding/hex" "fmt" "io" - "log" "os" "os/exec" "strings" + "sync" + "syscall" ) // Middleware represents a middleware object @@ -20,6 +21,8 @@ type Middleware struct { Stdout io.Reader commandCancel context.CancelFunc stop chan bool // Channel used only to indicate goroutine should shutdown + closed bool + mu sync.RWMutex } // NewMiddleware returns new middleware @@ -42,16 +45,19 @@ func NewMiddleware(command string) *Middleware { go m.read(m.Stdout) go func() { - err := cmd.Start() - - if err != nil { - log.Fatal(err) + defer m.Close() + var err error + if err = cmd.Start(); err == nil { + err = cmd.Wait() } - - err = cmd.Wait() - if err != nil { - log.Fatal(err) + if e, ok := err.(*exec.ExitError); ok { + status := e.Sys().(syscall.WaitStatus) + if status.Signal() == syscall.SIGKILL /*killed or context canceld */ { + return + } + } + Debug(0, fmt.Sprintf("[MIDDLEWARE] command[%q] error: %q", command, err.Error())) } }() @@ -60,7 +66,7 @@ func NewMiddleware(command string) *Middleware { // ReadFrom start a worker to read from this plugin func (m *Middleware) ReadFrom(plugin PluginReader) { - Debug(2, "[MIDDLEWARE-MASTER] Starting reading from", plugin) + Debug(2, fmt.Sprintf("[MIDDLEWARE] command[%q] Starting reading from %q", m.command, plugin)) go m.copy(m.Stdin, plugin) } @@ -75,17 +81,26 @@ func (m *Middleware) copy(to io.Writer, from PluginReader) { if msg == nil || len(msg.Data) == 0 { continue } - buf = msg.Data if Settings.PrettifyHTTP { buf = prettifyHTTP(msg.Data) } - dst = make([]byte, len(buf)*2+1) - hex.Encode(dst, buf) - dst[len(buf)*2] = '\n' - - to.Write(dst) + dstLen := (len(buf)+len(msg.Meta))*2 + 1 + // if enough space was previously allocated use it instead + if dstLen > len(dst) { + dst = make([]byte, dstLen) + } + n := hex.Encode(dst, msg.Meta) + n += hex.Encode(dst[n:], buf) + dst[n] = '\n' + n, err = to.Write(dst[:n+1]) + if err == nil { + continue + } + if m.isClosed() { + return + } } } @@ -93,19 +108,16 @@ func (m *Middleware) read(from io.Reader) { reader := bufio.NewReader(from) var line []byte var e error - for { if line, e = reader.ReadBytes('\n'); e != nil { - if e == io.EOF { - continue - } else { - break + if m.isClosed() { + return } + continue } - - buf := make([]byte, len(line)/2-1) + buf := make([]byte, (len(line)-1)/2) if _, err := hex.Decode(buf, line[:len(line)-1]); err != nil { - Debug(0, fmt.Sprintf("[MIDDLEWARE] failed to decode err: %q", err)) + Debug(0, fmt.Sprintf("[MIDDLEWARE] command[%q] failed to decode err: %q", m.command, err)) continue } var msg Message @@ -117,7 +129,6 @@ func (m *Middleware) read(from io.Reader) { } } - return } // PluginRead reads message from this plugin @@ -135,9 +146,21 @@ func (m *Middleware) String() string { return fmt.Sprintf("Modifying traffic using %q command", m.command) } +func (m *Middleware) isClosed() bool { + m.mu.RLock() + defer m.mu.RUnlock() + return m.closed +} + // Close closes this plugin func (m *Middleware) Close() error { + if m.isClosed() { + return nil + } + m.mu.Lock() + defer m.mu.Unlock() m.commandCancel() close(m.stop) + m.closed = true return nil } diff --git a/middleware_test.go b/middleware_test.go index c5ed1b3..c5884e7 100644 --- a/middleware_test.go +++ b/middleware_test.go @@ -2,225 +2,153 @@ package main import ( "bytes" - "crypto/rand" - "encoding/hex" - "net/http" - "net/http/httptest" - "net/http/httputil" + "context" + "os" + "os/exec" "strings" - "sync" + "sync/atomic" + "syscall" "testing" - "time" - "github.com/buger/goreplay/capture" "github.com/buger/goreplay/proto" ) -type fakeServiceCb func(string, int, []byte) +const echoSh = "./examples/middleware/echo.sh" +const tokenModifier = "go run ./examples/middleware/token_modifier.go" -// Simple service that generate token on request, and require this token for accesing to secure area -func NewFakeSecureService(wg *sync.WaitGroup, cb fakeServiceCb) *httptest.Server { - activeTokens := make([]string, 0) - var mu sync.Mutex +var noDebug = append(syscall.Environ(), "GOR_TEST=1") - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { - mu.Lock() - defer mu.Unlock() +func initMiddleware(cmd *exec.Cmd, cancl context.CancelFunc, l PluginReader, c func(error)) *Middleware { + var m Middleware + m.data = make(chan *Message, 1000) + m.stop = make(chan bool) + m.commandCancel = cancl + m.Stdout, _ = cmd.StdoutPipe() + m.Stdin, _ = cmd.StdinPipe() + cmd.Stderr = os.Stderr + go m.read(m.Stdout) + go func() { + defer m.Close() + var err error + if err = cmd.Start(); err == nil { + err = cmd.Wait() + } + if err != nil { + c(err) + } + }() + m.ReadFrom(l) + return &m +} - switch req.URL.Path { - case "/token": - // Generate random token - tokenLength := 10 - buf := make([]byte, tokenLength) - rand.Read(buf) - token := hex.EncodeToString(buf) - activeTokens = append(activeTokens, token) +func initCmd(command string, env []string) (*exec.Cmd, context.CancelFunc) { + commands := strings.Split(command, " ") + ctx, cancl := context.WithCancel(context.Background()) + cmd := exec.CommandContext(ctx, commands[0], commands[1:]...) + cmd.Env = env + return cmd, cancl +} - w.Write([]byte(token)) - - cb(req.URL.Path, 200, []byte(token)) - case "/secure": - token := req.URL.Query().Get("token") - tokenFound := false - - for _, t := range activeTokens { - if t == token { - tokenFound = true - break +func TestMiddlewareEarlyClose(t *testing.T) { + quit := make(chan struct{}) + in := NewTestInput() + cmd, cancl := initCmd(echoSh, noDebug) + midd := initMiddleware(cmd, cancl, in, func(err error) { + if err != nil { + if e, ok := err.(*exec.ExitError); ok { + status := e.Sys().(syscall.WaitStatus) + if status.Signal() != syscall.SIGKILL { + t.Errorf("expected error to be signal killed. got %s", status.Signal().String()) } } - - if tokenFound { - w.WriteHeader(http.StatusAccepted) - cb(req.URL.Path, 202, nil) - } else { - w.WriteHeader(http.StatusForbidden) - cb(req.URL.Path, 403, nil) - } } - - wg.Done() - })) - - return server -} - -func TestFakeSecureService(t *testing.T) { - var resp, token []byte - - wg := new(sync.WaitGroup) - - server := NewFakeSecureService(wg, func(path string, status int, resp []byte) { + quit <- struct{}{} }) - defer server.Close() - - wg.Add(3) - - client := NewHTTPClient(&HTTPOutputConfig{}).Client - rep, _ := client.Get(server.URL + "/token") - resp, _ = httputil.DumpResponse(rep, true) - token = proto.Body(resp) - - // Right token - rep, _ = client.Get(server.URL + "/secure?token=" + string(token)) - resp, _ = httputil.DumpResponse(rep, true) - if !bytes.Equal(proto.Status(resp), []byte("202")) { - t.Error("Valid token should return status 202:", string(proto.Status(resp))) + var body = []byte("OPTIONS / HTTP/1.1\r\nHost: example.org\r\n\r\n") + count := uint32(0) + out := NewTestOutput(func(msg *Message) { + if !bytes.Equal(body, msg.Data) { + t.Errorf("expected %q to equal %q", body, msg.Data) + } + atomic.AddUint32(&count, 1) + if atomic.LoadUint32(&count) == 5 { + quit <- struct{}{} + } + }) + pl := &InOutPlugins{} + pl.Inputs = []PluginReader{midd, in} + pl.Outputs = []PluginWriter{out} + pl.All = []interface{}{midd, out, in} + e := NewEmitter() + go e.Start(pl, "") + for i := 0; i < 5; i++ { + in.EmitBytes(body) } - - // Wrong tokens forbidden - rep, _ = client.Get(server.URL + "/secure?token=wrong") - resp, _ = httputil.DumpResponse(rep, true) - if !bytes.Equal(proto.Status(resp), []byte("403")) { - t.Error("Wrong token should returns status 403:", string(proto.Status(resp))) - } - - wg.Wait() -} - -func TestEchoMiddleware(t *testing.T) { - wg := new(sync.WaitGroup) - - from := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Env", "prod") - w.Header().Set("RequestPath", r.URL.Path) - wg.Done() - })) - defer from.Close() - - to := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Env", "test") - w.Header().Set("RequestPath", r.URL.Path) - wg.Done() - })) - defer to.Close() - - // Catch traffic from one service - fromAddr := strings.Replace(from.Listener.Addr().String(), "[::]", "127.0.0.1", -1) - conf := RAWInputConfig{ - Engine: capture.EnginePcap, - Expire: testRawExpire, - Protocol: ProtocolHTTP, - TrackResponse: true, - } - input := NewRAWInput(fromAddr, conf) - - // And redirect to another - output := NewHTTPOutput(to.URL, &HTTPOutputConfig{}) - - plugins := &InOutPlugins{ - Inputs: []PluginReader{input}, - Outputs: []PluginWriter{output}, - } - plugins.All = append(plugins.All, input, output) - - // Start Gor - emitter := NewEmitter() - emitter.Start(plugins, "echo -n && GOR_TEST=true && ./examples/middleware/echo.sh") - - // Wait till middleware initialization - time.Sleep(100 * time.Millisecond) - - // Should receive 2 requests from original + 2 from replayed - client := NewHTTPClient(output.(*HTTPOutput).config).Client - - for i := 0; i < 10; i++ { - wg.Add(2) - // Request should be echoed - client.Get(to.URL + "/a") - time.Sleep(5 * time.Millisecond) - client.Get(to.URL + "/b") - time.Sleep(5 * time.Millisecond) - } - - wg.Wait() - emitter.Close() + <-quit + midd.Close() + <-quit } func TestTokenMiddleware(t *testing.T) { - var resp, token []byte - - wg := new(sync.WaitGroup) - - from := NewFakeSecureService(wg, func(path string, status int, tok []byte) { - time.Sleep(10 * time.Millisecond) - }) - defer from.Close() - - to := NewFakeSecureService(wg, func(path string, status int, tok []byte) { - switch path { - case "/secure": - if status != 202 { - t.Error("Server should receive valid rewritten token") + quit := make(chan struct{}) + in := NewTestInput() + in.skipHeader = true + cmd, cancl := initCmd(tokenModifier, noDebug) + midd := initMiddleware(cmd, cancl, in, func(err error) {}) + req := []byte("1 932079936fa4306fc308d67588178d17d823647c 1439818823587396305 200\nGET /token HTTP/1.1\r\nHost: example.org\r\n\r\n") + res := []byte("2 932079936fa4306fc308d67588178d17d823647c 1439818823587396305 200\nHTTP/1.1 200 OK\r\nContent-Length: 10\r\nContent-Type: text/plain; charset=utf-8\r\n\r\n17d823647c") + rep := []byte("3 932079936fa4306fc308d67588178d17d823647c 1439818823587396305 200\nHTTP/1.1 200 OK\r\nContent-Length: 15\r\nContent-Type: text/plain; charset=utf-8\r\n\r\n932079936fa4306") + count := uint32(0) + out := NewTestOutput(func(msg *Message) { + if msg.Meta[0] == '1' && !bytes.Equal(payloadID(msg.Meta), payloadID(req)) { + token, _, _ := proto.PathParam(msg.Data, []byte("token")) + if !bytes.Equal(token, proto.Body(rep)) { + t.Error("expected the token to be equal to the replayed responses's token") } } - - time.Sleep(10 * time.Millisecond) + atomic.AddUint32(&count, 1) + if atomic.LoadUint32(&count) == 2 { + quit <- struct{}{} + } }) - defer to.Close() - - Settings.Middleware = "echo -n && GOR_TEST=true && go run ./examples/middleware/token_modifier.go" - - fromAddr := strings.Replace(from.Listener.Addr().String(), "[::]", "127.0.0.1", -1) - conf := RAWInputConfig{ - Engine: capture.EnginePcap, - Expire: testRawExpire, - Protocol: ProtocolHTTP, - TrackResponse: true, - } - // Catch traffic from one service - input := NewRAWInput(fromAddr, conf) - - // And redirect to another - output := NewHTTPOutput(to.URL, &HTTPOutputConfig{}) - - plugins := &InOutPlugins{ - Inputs: []PluginReader{input}, - Outputs: []PluginWriter{output}, - } - plugins.All = append(plugins.All, input, output) - - // Start Gor - emitter := NewEmitter() - emitter.Start(plugins, Settings.Middleware) - - // Should receive 2 requests from original + 2 from replayed - wg.Add(2) - - client := NewHTTPClient(&HTTPOutputConfig{}).Client - - // Sending traffic to original service - rep, _ := client.Get(to.URL + "/token") - resp, _ = httputil.DumpResponse(rep, true) - token = proto.Body(resp) - - rep, _ = client.Get(to.URL + "/secure?token=" + string(token)) - resp, _ = httputil.DumpResponse(rep, true) - if !bytes.Equal(proto.Status(resp), []byte("202")) { - t.Error("Valid token should return 202:", proto.Status(resp)) - } - - wg.Wait() - emitter.Close() - Settings.Middleware = "" + pl := &InOutPlugins{} + pl.Inputs = []PluginReader{midd, in} + pl.Outputs = []PluginWriter{out} + pl.All = []interface{}{midd, out, in} + e := NewEmitter() + go e.Start(pl, "") + in.EmitBytes(req) // emit original request + in.EmitBytes(res) // emit its response + in.EmitBytes(rep) // emit replayed response + // emit the request which should have modified token + token := []byte("1 8e091765ae902fef8a2b7d9dd96 14398188235873 100\nGET /?token=17d823647c HTTP/1.1\r\nHost: example.org\r\n\r\n") + in.EmitBytes(token) + <-quit + midd.Close() +} + +func TestMiddlewareWithPrettify(t *testing.T) { + Settings.PrettifyHTTP = true + quit := make(chan struct{}) + in := NewTestInput() + cmd, cancl := initCmd(echoSh, noDebug) + midd := initMiddleware(cmd, cancl, in, func(err error) {}) + var b1 = []byte("POST / HTTP/1.1\r\nHost: example.org\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nWiki\r\n5\r\npedia\r\nE\r\n in\r\n\r\nchunks.\r\n0\r\n\r\n") + var b2 = []byte("POST / HTTP/1.1\r\nHost: example.org\r\nContent-Length: 25\r\n\r\nWikipedia in\r\n\r\nchunks.") + out := NewTestOutput(func(msg *Message) { + if !bytes.Equal(proto.Body(b2), proto.Body(msg.Data)) { + t.Errorf("expected %q body to equal %q body", b2, msg.Data) + } + quit <- struct{}{} + }) + pl := &InOutPlugins{} + pl.Inputs = []PluginReader{midd, in} + pl.Outputs = []PluginWriter{out} + pl.All = []interface{}{midd, out, in} + e := NewEmitter() + go e.Start(pl, "") + in.EmitBytes(b1) + <-quit + midd.Close() + Settings.PrettifyHTTP = false }