diff --git a/.gitignore b/.gitignore index 2c6c95f..b03350b 100644 --- a/.gitignore +++ b/.gitignore @@ -5,8 +5,9 @@ *.deb *.pkg *.exe - +*.pprof *.out +hey *.bin diff --git a/capture/capture.go b/capture/capture.go index ab99ddc..b991e0b 100644 --- a/capture/capture.go +++ b/capture/capture.go @@ -209,6 +209,8 @@ func (l *Listener) Filter(ifi pcap.Interface) (filter string) { filter = fmt.Sprintf("%s or %s", filter, responseFilter) } + // filter = fmt.Sprintf("((((ip[2:2] - ((ip[0]&0xf)<<2)) - ((tcp[12]&0xf0)>>2)) != 0)) and (%s)", filter) + return } @@ -221,6 +223,9 @@ func (l *Listener) PcapHandle(ifi pcap.Interface) (handle *pcap.Handle, err erro return nil, fmt.Errorf("inactive handle error: %q, interface: %q", err, ifi.Name) } defer inactive.CleanUp() + + inactive.SetTimeout(pcap.BlockForever) + if l.TimestampType != "" { var ts pcap.TimestampSource ts, err = pcap.TimestampSourceFromString(l.TimestampType) @@ -364,11 +369,10 @@ func (l *Listener) closeHandles(key string) { l.Lock() defer l.Unlock() if handle, ok := l.Handles[key]; ok { - if _, ok = handle.(Socket); ok { - handle.(Socket).Close() - } else { - handle.(*pcap.Handle).Close() + if c, ok := handle.(io.Closer); ok { + c.Close() } + delete(l.Handles, key) if len(l.Handles) == 0 { close(l.closeDone) diff --git a/input_raw.go b/input_raw.go index aa6c7c3..5c6eb00 100644 --- a/input_raw.go +++ b/input_raw.go @@ -71,7 +71,7 @@ type RAWInput struct { RAWInputConfig messageStats []tcp.Stats listener *capture.Listener - message chan *tcp.Message + messageParser *tcp.MessageParser cancelListener context.CancelFunc closed bool } @@ -80,7 +80,6 @@ type RAWInput struct { func NewRAWInput(address string, config RAWInputConfig) (i *RAWInput) { i = new(RAWInput) i.RAWInputConfig = config - i.message = make(chan *tcp.Message, 10000) i.quit = make(chan bool) host, _ports, err := net.SplitHostPort(address) @@ -117,9 +116,11 @@ func (i *RAWInput) PluginRead() (*Message, error) { select { case <-i.quit: return nil, ErrorStopped - case msgTCP = <-i.message: + default: + msgTCP = i.messageParser.Read() msg.Data = msgTCP.Data() } + var msgType byte = ResponsePayload if msgTCP.IsRequest { msgType = RequestPayload @@ -157,15 +158,15 @@ func (i *RAWInput) listen(address string) { if err != nil { log.Fatal(err) } - parser := tcp.NewMessageParser(i.CopyBufferSize, i.Expire, Debug, i.messageEmitter) + i.messageParser = tcp.NewMessageParser(i.CopyBufferSize, i.Expire, Debug) if i.Protocol == ProtocolHTTP { - parser.Start = http1StartHint - parser.End = http1EndHint + i.messageParser.Start = http1StartHint + i.messageParser.End = http1EndHint } var ctx context.Context ctx, i.cancelListener = context.WithCancel(context.Background()) - errCh := i.listener.ListenBackground(ctx, parser.PacketHandler) + errCh := i.listener.ListenBackground(ctx, i.messageParser.PacketHandler) <-i.listener.Reading Debug(1, i) go func() { @@ -174,10 +175,6 @@ func (i *RAWInput) listen(address string) { }() } -func (i *RAWInput) messageEmitter(m *tcp.Message) { - i.message <- m -} - func (i *RAWInput) String() string { return fmt.Sprintf("Intercepting traffic from: %s:%s", i.host, strings.Join(strings.Fields(fmt.Sprint(i.ports)), ",")) } diff --git a/ring/ring.go b/ring/ring.go new file mode 100644 index 0000000..638277c --- /dev/null +++ b/ring/ring.go @@ -0,0 +1,205 @@ +package ring + +import ( + "errors" + "runtime" + "sync/atomic" + "time" +) + +var ( + // ErrDisposed is returned when an operation is performed on a disposed + // queue. + ErrDisposed = errors.New(`queue: disposed`) + + // ErrTimeout is returned when an applicable queue operation times out. + ErrTimeout = errors.New(`queue: poll timed out`) + + // ErrEmptyQueue is returned when an non-applicable queue operation was called + // due to the queue's empty item state + ErrEmptyQueue = errors.New(`queue: empty queue`) +) + +// roundUp takes a uint64 greater than 0 and rounds it up to the next +// power of 2. +func roundUp(v uint64) uint64 { + v-- + v |= v >> 1 + v |= v >> 2 + v |= v >> 4 + v |= v >> 8 + v |= v >> 16 + v |= v >> 32 + v++ + return v +} + +type node struct { + position uint64 + data interface{} +} + +type nodes []node + +// RingBuffer is a MPMC buffer that achieves threadsafety with CAS operations +// only. A put on full or get on empty call will block until an item +// is put or retrieved. Calling Dispose on the RingBuffer will unblock +// any blocked threads with an error. This buffer is similar to the buffer +// described here: http://www.1024cores.net/home/lock-free-algorithms/queues/bounded-mpmc-queue +// with some minor additions. +type RingBuffer struct { + _padding0 [8]uint64 + queue uint64 + _padding1 [8]uint64 + dequeue uint64 + _padding2 [8]uint64 + mask, disposed uint64 + _padding3 [8]uint64 + nodes nodes +} + +func (rb *RingBuffer) init(size uint64) { + size = roundUp(size) + rb.nodes = make(nodes, size) + for i := uint64(0); i < size; i++ { + rb.nodes[i] = node{position: i} + } + rb.mask = size - 1 // so we don't have to do this with every put/get operation +} + +// Put adds the provided item to the queue. If the queue is full, this +// call will block until an item is added to the queue or Dispose is called +// on the queue. An error will be returned if the queue is disposed. +func (rb *RingBuffer) Put(item interface{}) error { + _, err := rb.put(item, false) + return err +} + +// Offer adds the provided item to the queue if there is space. If the queue +// is full, this call will return false. An error will be returned if the +// queue is disposed. +func (rb *RingBuffer) Offer(item interface{}) (bool, error) { + return rb.put(item, true) +} + +func (rb *RingBuffer) put(item interface{}, offer bool) (bool, error) { + var n *node + pos := atomic.LoadUint64(&rb.queue) +L: + for { + if atomic.LoadUint64(&rb.disposed) == 1 { + return false, ErrDisposed + } + + n = &rb.nodes[pos&rb.mask] + seq := atomic.LoadUint64(&n.position) + switch dif := seq - pos; { + case dif == 0: + if atomic.CompareAndSwapUint64(&rb.queue, pos, pos+1) { + break L + } + case dif < 0: + panic(`Ring buffer in a compromised state during a put operation.`) + default: + pos = atomic.LoadUint64(&rb.queue) + } + + if offer { + return false, nil + } + + runtime.Gosched() // free up the cpu before the next iteration + } + + n.data = item + atomic.StoreUint64(&n.position, pos+1) + return true, nil +} + +// Get will return the next item in the queue. This call will block +// if the queue is empty. This call will unblock when an item is added +// to the queue or Dispose is called on the queue. An error will be returned +// if the queue is disposed. +func (rb *RingBuffer) Get() (interface{}, error) { + return rb.Poll(0) +} + +// Poll will return the next item in the queue. This call will block +// if the queue is empty. This call will unblock when an item is added +// to the queue, Dispose is called on the queue, or the timeout is reached. An +// error will be returned if the queue is disposed or a timeout occurs. A +// non-positive timeout will block indefinitely. +func (rb *RingBuffer) Poll(timeout time.Duration) (interface{}, error) { + var ( + n *node + pos = atomic.LoadUint64(&rb.dequeue) + start time.Time + ) + if timeout > 0 { + start = time.Now() + } +L: + for { + if atomic.LoadUint64(&rb.disposed) == 1 { + return nil, ErrDisposed + } + + n = &rb.nodes[pos&rb.mask] + seq := atomic.LoadUint64(&n.position) + switch dif := seq - (pos + 1); { + case dif == 0: + if atomic.CompareAndSwapUint64(&rb.dequeue, pos, pos+1) { + break L + } + case dif < 0: + panic(`Ring buffer in compromised state during a get operation.`) + default: + pos = atomic.LoadUint64(&rb.dequeue) + } + + if timeout > 0 && time.Since(start) >= timeout { + return nil, ErrTimeout + } + + if timeout < 0 { + return nil, ErrTimeout + } + + runtime.Gosched() // free up the cpu before the next iteration + } + data := n.data + n.data = nil + atomic.StoreUint64(&n.position, pos+rb.mask+1) + return data, nil +} + +// Len returns the number of items in the queue. +func (rb *RingBuffer) Len() uint64 { + return atomic.LoadUint64(&rb.queue) - atomic.LoadUint64(&rb.dequeue) +} + +// Cap returns the capacity of this ring buffer. +func (rb *RingBuffer) Cap() uint64 { + return uint64(len(rb.nodes)) +} + +// Dispose will dispose of this queue and free any blocked threads +// in the Put and/or Get methods. Calling those methods on a disposed +// queue will return an error. +func (rb *RingBuffer) Dispose() { + atomic.CompareAndSwapUint64(&rb.disposed, 0, 1) +} + +// IsDisposed will return a bool indicating if this queue has been +// disposed. +func (rb *RingBuffer) IsDisposed() bool { + return atomic.LoadUint64(&rb.disposed) == 1 +} + +// NewRingBuffer will allocate, initialize, and return a ring buffer +// with the specified size. +func NewRingBuffer(size uint64) *RingBuffer { + rb := &RingBuffer{} + rb.init(size) + return rb +} diff --git a/tcp/tcp_message.go b/tcp/tcp_message.go index b45aa05..e10a9ee 100644 --- a/tcp/tcp_message.go +++ b/tcp/tcp_message.go @@ -7,6 +7,7 @@ import ( "sort" "time" + "github.com/buger/goreplay/ring" "github.com/buger/goreplay/size" ) @@ -178,24 +179,23 @@ type HintStart func(*Packet) (IsRequest, IsOutgoing bool) // MessageParser holds data of all tcp messages in progress(still receiving/sending packets). // message is identified by its source port and dst port, and last 4bytes of src IP. type MessageParser struct { - debug Debugger - maxSize size.Size // maximum message size, default 5mb - m map[uint64]*Message - emit Emitter + debug Debugger + maxSize size.Size // maximum message size, default 5mb + m map[uint64]*Message + messageExpire time.Duration // the maximum time to wait for the final packet, minimum is 100ms End HintEnd Start HintStart ticker *time.Ticker - packets chan *Packet - msgs int32 // messages in the parser + messages *ring.RingBuffer + packets *ring.RingBuffer close chan struct{} // to signal that we are able to close } // NewMessageParser returns a new instance of message parser -func NewMessageParser(maxSize size.Size, messageExpire time.Duration, debugger Debugger, emitHandler Emitter) (parser *MessageParser) { +func NewMessageParser(maxSize size.Size, messageExpire time.Duration, debugger Debugger) (parser *MessageParser) { parser = new(MessageParser) parser.debug = debugger - parser.emit = emitHandler parser.messageExpire = time.Millisecond * 100 if parser.messageExpire < messageExpire { parser.messageExpire = messageExpire @@ -204,7 +204,10 @@ func NewMessageParser(maxSize size.Size, messageExpire time.Duration, debugger D if parser.maxSize < 1 { parser.maxSize = 5 << 20 } - parser.packets = make(chan *Packet, 1000) + + parser.packets = ring.NewRingBuffer(10000) + parser.messages = ring.NewRingBuffer(10000) + parser.m = make(map[uint64]*Message) parser.ticker = time.NewTicker(time.Millisecond * 50) parser.close = make(chan struct{}, 1) @@ -214,18 +217,22 @@ func NewMessageParser(maxSize size.Size, messageExpire time.Duration, debugger D // Packet returns packet handler func (parser *MessageParser) PacketHandler(packet *Packet) { - parser.packets <- packet + parser.packets.Offer(packet) } func (parser *MessageParser) wait() { var ( - pckt *Packet - now time.Time + now time.Time ) for { + pckt, err := parser.packets.Poll(-1) + if err == nil { + parser.processPacket(pckt.(*Packet)) + } else { + time.Sleep(50 * time.Millisecond) + } + select { - case pckt = <-parser.packets: - parser.processPacket(pckt) case now = <-parser.ticker.C: parser.timer(now) case <-parser.close: @@ -254,11 +261,9 @@ func (parser *MessageParser) processPacket(pckt *Packet) { // Requeue not known packets pckt.Retry++ - select { - case parser.packets <- pckt: - default: + if ok, _ := parser.packets.Offer(pckt); !ok { + // Drop packet if it does not fit to ring buffer packetPool.Put(pckt) - // fmt.Println("Skipping packet") } } return @@ -292,9 +297,20 @@ func (parser *MessageParser) addPacket(m *Message, pckt *Packet) { parser.Emit(m) } +func (parser *MessageParser) Read() *Message { + for { + if m, err := parser.messages.Poll(-1); err != nil { + time.Sleep(50 * time.Millisecond) + } else { + return m.(*Message) + } + } +} + func (parser *MessageParser) Emit(m *Message) { delete(parser.m, m.packets[0].MessageID()) - parser.emit(m) + + parser.messages.Offer(m) } func (parser *MessageParser) timer(now time.Time) { diff --git a/tcp/tcp_packet.go b/tcp/tcp_packet.go index 21ecfc5..822aa2a 100644 --- a/tcp/tcp_packet.go +++ b/tcp/tcp_packet.go @@ -123,7 +123,7 @@ func ParsePacket(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo) (pc } if len(ndata[dOf:]) == 0 { - return nil, fmt.Errorf("Packet without Data") + return nil, EmptyPacket("") } if (netLayer[0] >> 4) == 4 { @@ -150,7 +150,7 @@ func ParsePacket(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo) (pc pckt.RST = transLayer[13]&0x04 != 0 pckt.ACK = transLayer[13]&0x10 != 0 pckt.Lost = uint32(cp.Length - cp.CaptureLength) - pckt.Payload = copySlice(pckt.Payload, ndata[dOf:]) + pckt.Payload = ndata[dOf:] return } @@ -174,6 +174,12 @@ func (pckt *Packet) Dst() string { return fmt.Sprintf("%s:%d", pckt.DstIP, pckt.DstPort) } +type EmptyPacket string + +func (err EmptyPacket) Error() string { + return "Empty packet" +} + // ErrHdrLength returned on short header length type ErrHdrLength string diff --git a/tcp/tcp_test.go b/tcp/tcp_test.go index 2dd4424..b066ed7 100644 --- a/tcp/tcp_test.go +++ b/tcp/tcp_test.go @@ -74,8 +74,7 @@ func TestRequestResponseMapping(t *testing.T) { {SrcPort: 80, DstPort: 60000, Ack: 71, Seq: 56, Timestamp: time.Unix(8, 0), Payload: []byte("Content-Length: 0\r\n\r\n")}, } - var mssg = make(chan *Message, 4) - parser := NewMessageParser(1<<20, time.Second, nil, func(m *Message) { mssg <- m }) + parser := NewMessageParser(1<<20, time.Second, nil) parser.Start = func(pckt *Packet) (bool, bool) { return proto.HasRequestTitle(pckt.Payload), proto.HasResponseTitle(pckt.Payload) } @@ -89,13 +88,8 @@ func TestRequestResponseMapping(t *testing.T) { messages := []*Message{} for i := 0; i < 4; i++ { - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m := <-mssg: - messages = append(messages, m) - } + m := parser.Read() + messages = append(messages, m) } assert.Equal(t, messages[0].IsRequest, true) @@ -110,8 +104,7 @@ func TestRequestResponseMapping(t *testing.T) { } func TestMessageParserWithHint(t *testing.T) { - var mssg = make(chan *Message, 3) - parser := NewMessageParser(1<<20, time.Second, nil, func(m *Message) { mssg <- m }) + parser := NewMessageParser(1<<20, time.Second, nil) parser.Start = func(pckt *Packet) (bool, bool) { return proto.HasRequestTitle(pckt.Payload), proto.HasResponseTitle(pckt.Payload) } @@ -133,13 +126,8 @@ func TestMessageParserWithHint(t *testing.T) { messages := []*Message{} for i := 0; i < 3; i++ { - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m := <-mssg: - messages = append(messages, m) - } + m := parser.Read() + messages = append(messages, m) } if !bytes.HasSuffix(messages[0].Data(), []byte("\n7\r\nNetwork\r\n0\r\n\r\n")) { @@ -156,8 +144,7 @@ func TestMessageParserWithHint(t *testing.T) { } func TestMessageParserWrongOrder(t *testing.T) { - var mssg = make(chan *Message, 3) - parser := NewMessageParser(1<<20, time.Second, nil, func(m *Message) { mssg <- m }) + parser := NewMessageParser(1<<20, time.Second, nil) parser.Start = func(pckt *Packet) (bool, bool) { return proto.HasRequestTitle(pckt.Payload), proto.HasResponseTitle(pckt.Payload) } @@ -178,66 +165,45 @@ func TestMessageParserWrongOrder(t *testing.T) { for i := 0; i < 30; i++ { parser.PacketHandler(packets[i]) } - var m *Message - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m = <-mssg: - } + + m := parser.Read() + if !bytes.HasSuffix(m.Data(), []byte("\n7\r\nNetwork\r\n0\r\n\r\n")) { t.Errorf("expected to %q to have suffix %q", m.Data(), []byte("\n7\r\nNetwork\r\n0\r\n\r\n")) } - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m = <-mssg: - } + m = parser.Read() + if !bytes.HasSuffix(m.Data(), []byte("Network")) { t.Errorf("expected to %q to have suffix %q", m.Data(), []byte("Network")) } } func TestMessageParserWithoutHint(t *testing.T) { - var mssg = make(chan *Message, 1) var data [63 << 10]byte packets := GetPackets(true, 1, 10, data[:]) - p := NewMessageParser(63<<10*10, time.Second, nil, func(m *Message) { mssg <- m }) + p := NewMessageParser(63<<10*10, time.Second, nil) for _, v := range packets { p.PacketHandler(v) } - var m *Message - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m = <-mssg: - } + m := p.Read() + if m.Length != 63<<10*10 { t.Errorf("expected %d to equal %d", m.Length, 63<<10*10) } } func TestMessageMaxSizeReached(t *testing.T) { - var mssg = make(chan *Message, 2) var data [63 << 10]byte packets := GetPackets(true, 1, 2, data[:]) packets = append(packets, GetPackets(true, 1, 1, make([]byte, 63<<10+10))...) - p := NewMessageParser(63<<10+10, time.Second, nil, func(m *Message) { mssg <- m }) + p := NewMessageParser(63<<10+10, time.Second, nil) for _, v := range packets { p.PacketHandler(v) } - var m *Message - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m = <-mssg: - } + m := p.Read() if m.Length != 63<<10+10 { t.Errorf("expected %d to equal %d", m.Length, 63<<10+10) } @@ -245,12 +211,8 @@ func TestMessageMaxSizeReached(t *testing.T) { t.Error("expected message to be truncated") } - select { - case <-time.After(time.Second): - t.Errorf("can't parse packets fast enough") - return - case m = <-mssg: - } + m = p.Read() + if m.Length != 63<<10+10 { t.Errorf("expected %d to equal %d", m.Length, 63<<10+10) } @@ -260,14 +222,13 @@ func TestMessageMaxSizeReached(t *testing.T) { } func TestMessageTimeoutReached(t *testing.T) { - var mssg = make(chan *Message, 2) var data [63 << 10]byte packets := GetPackets(true, 1, 2, data[:]) - p := NewMessageParser(1<<20, 0, nil, func(m *Message) { mssg <- m }) + p := NewMessageParser(1<<20, 0, nil) p.PacketHandler(packets[0]) time.Sleep(time.Millisecond * 400) p.PacketHandler(packets[1]) - m := <-mssg + m := p.Read() if m.Length != 63<<10 { t.Errorf("expected %d to equal %d", m.Length, 63<<10) } @@ -280,18 +241,18 @@ func TestMessageUUID(t *testing.T) { packets := GetPackets(true, 1, 10, nil) var uuid, uuid1 []byte - parser := NewMessageParser(0, 0, nil, func(msg *Message) { - if len(uuid) == 0 { - uuid = msg.UUID() - return - } - uuid1 = msg.UUID() - }) + parser := NewMessageParser(0, 0, nil) for _, p := range packets { parser.PacketHandler(p) } + m := parser.Read() + uuid = m.UUID() + + m = parser.Read() + uuid1 = m.UUID() + if string(uuid) != string(uuid1) { t.Errorf("expected %s, to equal %s", uuid, uuid1) } @@ -301,13 +262,13 @@ func BenchmarkMessageUUID(b *testing.B) { packets := GetPackets(true, 1, 5, nil) var uuid []byte - var msg *Message - parser := NewMessageParser(0, 0, nil, func(m *Message) { - msg = m - }) + parser := NewMessageParser(0, 0, nil) for _, p := range packets { parser.PacketHandler(p) } + + msg := parser.Read() + b.ResetTimer() for i := 0; i < b.N; i++ { uuid = msg.UUID() @@ -328,20 +289,16 @@ func BenchmarkPacketParseAndSort(b *testing.B) { } func BenchmarkMessageParserWithoutHint(b *testing.B) { - // runtime.GOMAXPROCS(8) - var mssg = make(chan *Message, 1) var chunk = []byte("111111111111111111111111111111") packets := GetPackets(true, 1, 1000, chunk) - p := NewMessageParser(1<<20, time.Second*2, nil, func(m *Message) { - mssg <- m - }) + p := NewMessageParser(1<<20, time.Second*2, nil) b.ResetTimer() b.ReportMetric(float64(1000), "packets/op") for i := 0; i < b.N; i++ { for _, v := range packets { p.PacketHandler(v) } - <-mssg + p.Read() } } @@ -357,8 +314,8 @@ func BenchmarkMessageParserWithHint(b *testing.B) { for i := 0; i < len(buf); i++ { packets[i] = GetPackets(false, 1, 1, buf[i])[0] } - var mssg = make(chan *Message, 1) - parser := NewMessageParser(1<<30, time.Second*10, nil, func(m *Message) { mssg <- m }) + + parser := NewMessageParser(1<<30, time.Second*10, nil) parser.Start = func(pckt *Packet) (bool, bool) { return false, proto.HasResponseTitle(pckt.Payload) } @@ -372,7 +329,7 @@ func BenchmarkMessageParserWithHint(b *testing.B) { for j := range packets { parser.PacketHandler(packets[j]) } - <-mssg + parser.Read() } }