diff --git a/capture/capture.go b/capture/capture.go index c438f12..5648c92 100644 --- a/capture/capture.go +++ b/capture/capture.go @@ -42,7 +42,7 @@ type Listener struct { sync.Mutex Transport string // transport layer default to tcp Activate func() error // function is used to activate the engine. it must be called before reading packets - Handles map[string]gopacket.ZeroCopyPacketDataSource + Handles map[string]packetHandle Interfaces []pcap.Interface loopIndex int Reading chan bool // this channel is closed when the listener has started reading packets @@ -57,6 +57,11 @@ type Listener struct { quit chan struct{} } +type packetHandle struct { + handler gopacket.ZeroCopyPacketDataSource + ips []net.IP +} + // EngineType ... type EngineType uint8 @@ -117,7 +122,7 @@ func NewListener(host string, ports []uint16, transport string, engine EngineTyp if transport != "" { l.Transport = transport } - l.Handles = make(map[string]gopacket.ZeroCopyPacketDataSource) + l.Handles = make(map[string]packetHandle) l.trackResponse = trackResponse l.closeDone = make(chan struct{}) l.quit = make(chan struct{}) @@ -312,12 +317,12 @@ func (l *Listener) read(handler PacketHandler) { l.Lock() defer l.Unlock() for key, handle := range l.Handles { - go func(key string, hndl gopacket.ZeroCopyPacketDataSource) { + go func(key string, hndl packetHandle) { defer l.closeHandles(key) linkSize := 14 linkType := int(layers.LinkTypeEthernet) - if _, ok := hndl.(*pcap.Handle); ok { - linkType = int(hndl.(*pcap.Handle).LinkType()) + if _, ok := hndl.handler.(*pcap.Handle); ok { + linkType = int(hndl.handler.(*pcap.Handle).LinkType()) linkSize, ok = pcapLinkTypeLength(linkType) if !ok { if os.Getenv("GORDEBUG") != "0" { @@ -332,10 +337,22 @@ func (l *Listener) read(handler PacketHandler) { case <-l.quit: return default: - data, ci, err := hndl.ZeroCopyReadPacketData() + data, ci, err := hndl.handler.ZeroCopyReadPacketData() if err == nil { - pckt, err := tcp.ParsePacket(data, linkType, linkSize, &ci) + pckt, err := tcp.ParsePacket(data, linkType, linkSize, &ci, false) if err == nil { + for _, p := range l.ports { + if pckt.DstPort == p { + for _, ip := range hndl.ips { + if pckt.DstIP.Equal(ip) { + pckt.Incoming = true + break + } + } + break + } + } + handler(pckt) } continue @@ -367,7 +384,7 @@ func (l *Listener) closeHandles(key string) { l.Lock() defer l.Unlock() if handle, ok := l.Handles[key]; ok { - if c, ok := handle.(io.Closer); ok { + if c, ok := handle.handler.(io.Closer); ok { c.Close() } @@ -388,7 +405,10 @@ func (l *Listener) activatePcap() error { msg += ("\n" + e.Error()) continue } - l.Handles[ifi.Name] = handle + l.Handles[ifi.Name] = packetHandle{ + handler: handle, + ips: interfaceIPs(ifi), + } } if len(l.Handles) == 0 { return fmt.Errorf("pcap handles error:%s", msg) @@ -409,7 +429,10 @@ func (l *Listener) activateRawSocket() error { msg += ("\n" + e.Error()) continue } - l.Handles[ifi.Name] = handle + l.Handles[ifi.Name] = packetHandle{ + handler: handle, + ips: interfaceIPs(ifi), + } } if len(l.Handles) == 0 { return fmt.Errorf("raw socket handles error:%s", msg) @@ -433,7 +456,9 @@ func (l *Listener) activatePcapFile() (err error) { handle.Close() return fmt.Errorf("BPF filter error: %q, filter: %s", e, l.BPFFilter) } - l.Handles["pcap_file"] = handle + l.Handles["pcap_file"] = packetHandle{ + handler: handle, + } return } @@ -456,7 +481,10 @@ func (l *Listener) activateAFPacket() error { fmt.Println("Interface:", ifi.Name, ". BPF Filter:", l.BPFFilter) handle.SetBPFFilter(l.BPFFilter, 64<<10) - l.Handles[ifi.Name] = handle + l.Handles[ifi.Name] = packetHandle{ + handler: handle, + ips: interfaceIPs(ifi), + } } if len(l.Handles) == 0 { @@ -524,6 +552,14 @@ func interfaceAddresses(ifi pcap.Interface) []string { return hosts } +func interfaceIPs(ifi pcap.Interface) []net.IP { + var ips []net.IP + for _, addr := range ifi.Addresses { + ips = append(ips, addr.IP) + } + return ips +} + func listenAll(addr string) bool { switch addr { case "", "0.0.0.0", "[::]", "::": diff --git a/input_raw.go b/input_raw.go index 5c6eb00..cec9a99 100644 --- a/input_raw.go +++ b/input_raw.go @@ -53,16 +53,17 @@ func (protocol *TCPProtocol) String() string { // RAWInputConfig represents configuration that can be applied on raw input type RAWInputConfig struct { capture.PcapOptions - Expire time.Duration `json:"input-raw-expire"` - CopyBufferSize size.Size `json:"copy-buffer-size"` - Engine capture.EngineType `json:"input-raw-engine"` - TrackResponse bool `json:"input-raw-track-response"` - Protocol TCPProtocol `json:"input-raw-protocol"` - RealIPHeader string `json:"input-raw-realip-header"` - Stats bool `json:"input-raw-stats"` - quit chan bool // Channel used only to indicate goroutine should shutdown - host string - ports []uint16 + Expire time.Duration `json:"input-raw-expire"` + CopyBufferSize size.Size `json:"copy-buffer-size"` + Engine capture.EngineType `json:"input-raw-engine"` + TrackResponse bool `json:"input-raw-track-response"` + Protocol TCPProtocol `json:"input-raw-protocol"` + RealIPHeader string `json:"input-raw-realip-header"` + Stats bool `json:"input-raw-stats"` + AllowIncomplete bool `json:"input-raw-allow-incomplete"` + quit chan bool // Channel used only to indicate goroutine should shutdown + host string + ports []uint16 } // RAWInput used for intercepting traffic for given address @@ -116,8 +117,7 @@ func (i *RAWInput) PluginRead() (*Message, error) { select { case <-i.quit: return nil, ErrorStopped - default: - msgTCP = i.messageParser.Read() + case msgTCP = <-i.messageParser.Messages(): msg.Data = msgTCP.Data() } @@ -158,7 +158,7 @@ func (i *RAWInput) listen(address string) { if err != nil { log.Fatal(err) } - i.messageParser = tcp.NewMessageParser(i.CopyBufferSize, i.Expire, Debug) + i.messageParser = tcp.NewMessageParser(i.CopyBufferSize, i.Expire, i.AllowIncomplete, Debug) if i.Protocol == ProtocolHTTP { i.messageParser.Start = http1StartHint diff --git a/input_raw_test.go b/input_raw_test.go index 9071396..d844f6a 100644 --- a/input_raw_test.go +++ b/input_raw_test.go @@ -58,6 +58,7 @@ func TestRAWInputIPv4(t *testing.T) { } else { respCounter++ } + wg.Done() }) @@ -71,14 +72,18 @@ func TestRAWInputIPv4(t *testing.T) { emitter := NewEmitter() defer emitter.Close() go emitter.Start(plugins, Settings.Middleware) + + // time.Sleep(time.Second) for i := 0; i < 1; i++ { wg.Add(2) _, err = http.Get(addr) + if err != nil { t.Error(err) return } } + wg.Wait() const want = 10 if reqCounter != respCounter && reqCounter != want { diff --git a/settings.go b/settings.go index 9a9cffa..954aa0a 100644 --- a/settings.go +++ b/settings.go @@ -141,6 +141,7 @@ func init() { flag.BoolVar(&Settings.Promiscuous, "input-raw-promisc", false, "enable promiscuous mode") flag.BoolVar(&Settings.Monitor, "input-raw-monitor", false, "enable RF monitor mode") flag.BoolVar(&Settings.Stats, "input-raw-stats", false, "enable stats generator on raw TCP messages") + flag.BoolVar(&Settings.AllowIncomplete, "input-raw-allow-incomplete", false, "If turned on Gor will record HTTP messages with missing packets") flag.StringVar(&Settings.Middleware, "middleware", "", "Used for modifying traffic using external command") diff --git a/tcp/tcp_message.go b/tcp/tcp_message.go index 4db342d..0ac1285 100644 --- a/tcp/tcp_message.go +++ b/tcp/tcp_message.go @@ -184,23 +184,27 @@ type MessageParser struct { 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 - messages chan *Message - packets chan *Packet - close chan struct{} // to signal that we are able to close + messageExpire time.Duration // the maximum time to wait for the final packet, minimum is 100ms + allowIncompete bool + End HintEnd + Start HintStart + ticker *time.Ticker + messages chan *Message + packets chan *Packet + 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) (parser *MessageParser) { +func NewMessageParser(maxSize size.Size, messageExpire time.Duration, allowIncompete bool, debugger Debugger) (parser *MessageParser) { parser = new(MessageParser) parser.debug = debugger - parser.messageExpire = time.Millisecond * 100 - if parser.messageExpire < messageExpire { - parser.messageExpire = messageExpire + + parser.messageExpire = messageExpire + if parser.messageExpire == 0 { + parser.messageExpire = time.Millisecond * 500 } + + parser.allowIncompete = allowIncompete parser.maxSize = maxSize if parser.maxSize < 1 { parser.maxSize = 5 << 20 @@ -224,8 +228,6 @@ func (parser *MessageParser) PacketHandler(packet *Packet) { parser.packets <- packet } -var processedPackets int - func (parser *MessageParser) wait() { var ( now time.Time @@ -271,6 +273,8 @@ func (parser *MessageParser) processPacket(pckt *Packet) { } return } + default: + in = pckt.Incoming } m = new(Message) @@ -288,16 +292,18 @@ func (parser *MessageParser) addPacket(m *Message, pckt *Packet) { pckt.Payload = pckt.Payload[:int(parser.maxSize)-m.Length] } m.add(pckt) - switch { - // if one of this cases matches, we dispatch the message - case trunc >= 0: - case parser.End != nil && parser.End(m): - default: - // continue to receive packets + + if trunc > 0 { return } - parser.Emit(m) + // If we are using protocol parsing, like HTTP, depend on its parsing func. + // For the binary procols wait for message to expire + if parser.End != nil { + if parser.End(m) { + parser.Emit(m) + } + } } func (parser *MessageParser) Read() *Message { @@ -305,6 +311,10 @@ func (parser *MessageParser) Read() *Message { return m } +func (parser *MessageParser) Messages() chan *Message { + return parser.messages +} + func (parser *MessageParser) Emit(m *Message) { delete(parser.m, m.packets[0].MessageID()) @@ -321,7 +331,13 @@ func (parser *MessageParser) timer(now time.Time) { for _, m := range parser.m { if now.Sub(m.End) > parser.messageExpire { m.TimedOut = true - parser.Emit(m) + if parser.End == nil || parser.allowIncompete { + parser.Emit(m) + } else { + // Just remove + delete(parser.m, m.packets[0].MessageID()) + m.Finalize() + } } } } diff --git a/tcp/tcp_packet.go b/tcp/tcp_packet.go index 7786c69..7740a2a 100644 --- a/tcp/tcp_packet.go +++ b/tcp/tcp_packet.go @@ -58,6 +58,7 @@ calllers must make sure that ParsePacket has'nt returned any error before callin function. */ type Packet struct { + Incoming bool messageID uint64 SrcIP, DstIP net.IP Version uint8 @@ -72,9 +73,9 @@ type Packet struct { } // ParsePacket parse raw packets -func ParsePacket(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo) (pckt *Packet, err error) { +func ParsePacket(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo, allowEmpty bool) (pckt *Packet, err error) { pckt = packetPool.Get() - if err := pckt.parse(data, lType, lTypeLen, cp); err != nil { + if err := pckt.parse(data, lType, lTypeLen, cp, allowEmpty); err != nil { packetPool.Put(pckt) return nil, err } @@ -82,7 +83,7 @@ func ParsePacket(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo) (pc return pckt, nil } -func (pckt *Packet) parse(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo) error { +func (pckt *Packet) parse(data []byte, lType, lTypeLen int, cp *gopacket.CaptureInfo, allowEmpty bool) error { pckt.Retry = 0 pckt.messageID = 0 @@ -158,7 +159,7 @@ func (pckt *Packet) parse(data []byte, lType, lTypeLen int, cp *gopacket.Capture return ErrHdrLength("TCP opts") } - if len(ndata[dOf:]) == 0 { + if !allowEmpty && len(ndata[dOf:]) == 0 { return EmptyPacket("") } diff --git a/tcp/tcp_test.go b/tcp/tcp_test.go index b066ed7..d05a631 100644 --- a/tcp/tcp_test.go +++ b/tcp/tcp_test.go @@ -43,14 +43,15 @@ func generateHeader(request bool, seq uint32, length uint16) []byte { func GetPackets(request bool, start uint32, _len int, payload []byte) []*Packet { var packets = make([]*Packet, _len) + var err error for i := start; i < start+uint32(_len); i++ { d := append(generateHeader(request, i, uint16(len(payload))), payload...) ci := &gopacket.CaptureInfo{Length: len(d), CaptureLength: len(d), Timestamp: time.Now()} - if len(payload) > 0 { - packets[i-start], _ = ParsePacket(d, int(layers.LinkTypeLoop), 4, ci) - } else { - packets[i-start] = new(Packet) + packets[i-start], err = ParsePacket(d, int(layers.LinkTypeLoop), 4, ci, true) + packets[i-start].Incoming = request + if err != nil { + panic(err) } } return packets @@ -74,7 +75,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")}, } - parser := NewMessageParser(1<<20, time.Second, nil) + parser := NewMessageParser(1<<20, time.Second, false, nil) parser.Start = func(pckt *Packet) (bool, bool) { return proto.HasRequestTitle(pckt.Payload), proto.HasResponseTitle(pckt.Payload) } @@ -104,7 +105,7 @@ func TestRequestResponseMapping(t *testing.T) { } func TestMessageParserWithHint(t *testing.T) { - parser := NewMessageParser(1<<20, time.Second, nil) + parser := NewMessageParser(1<<20, time.Second, false, nil) parser.Start = func(pckt *Packet) (bool, bool) { return proto.HasRequestTitle(pckt.Payload), proto.HasResponseTitle(pckt.Payload) } @@ -144,7 +145,7 @@ func TestMessageParserWithHint(t *testing.T) { } func TestMessageParserWrongOrder(t *testing.T) { - parser := NewMessageParser(1<<20, time.Second, nil) + parser := NewMessageParser(1<<20, time.Second, true, nil) parser.Start = func(pckt *Packet) (bool, bool) { return proto.HasRequestTitle(pckt.Payload), proto.HasResponseTitle(pckt.Payload) } @@ -183,7 +184,7 @@ func TestMessageParserWithoutHint(t *testing.T) { var data [63 << 10]byte packets := GetPackets(true, 1, 10, data[:]) - p := NewMessageParser(63<<10*10, time.Second, nil) + p := NewMessageParser(63<<10*10, time.Second, false, nil) for _, v := range packets { p.PacketHandler(v) } @@ -197,12 +198,14 @@ func TestMessageParserWithoutHint(t *testing.T) { func TestMessageMaxSizeReached(t *testing.T) { var data [63 << 10]byte packets := GetPackets(true, 1, 2, data[:]) - packets = append(packets, GetPackets(true, 1, 1, make([]byte, 63<<10+10))...) + packets = append(packets, GetPackets(false, 1, 1, make([]byte, 63<<10+10))...) - p := NewMessageParser(63<<10+10, time.Second, nil) + p := NewMessageParser(63<<10+10, time.Millisecond, false, nil) for _, v := range packets { p.PacketHandler(v) } + time.Sleep(10 * time.Millisecond) + m := p.Read() if m.Length != 63<<10+10 { t.Errorf("expected %d to equal %d", m.Length, 63<<10+10) @@ -224,9 +227,9 @@ func TestMessageMaxSizeReached(t *testing.T) { func TestMessageTimeoutReached(t *testing.T) { var data [63 << 10]byte packets := GetPackets(true, 1, 2, data[:]) - p := NewMessageParser(1<<20, 0, nil) + p := NewMessageParser(1<<20, 10*time.Millisecond, true, nil) p.PacketHandler(packets[0]) - time.Sleep(time.Millisecond * 400) + time.Sleep(time.Millisecond * 50) p.PacketHandler(packets[1]) m := p.Read() if m.Length != 63<<10 { @@ -237,32 +240,11 @@ func TestMessageTimeoutReached(t *testing.T) { } } -func TestMessageUUID(t *testing.T) { - packets := GetPackets(true, 1, 10, nil) - - var uuid, uuid1 []byte - 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) - } -} - func BenchmarkMessageUUID(b *testing.B) { packets := GetPackets(true, 1, 5, nil) var uuid []byte - parser := NewMessageParser(0, 0, nil) + parser := NewMessageParser(0, 0, false, nil) for _, p := range packets { parser.PacketHandler(p) } @@ -291,7 +273,7 @@ func BenchmarkPacketParseAndSort(b *testing.B) { func BenchmarkMessageParserWithoutHint(b *testing.B) { var chunk = []byte("111111111111111111111111111111") packets := GetPackets(true, 1, 1000, chunk) - p := NewMessageParser(1<<20, time.Second*2, nil) + p := NewMessageParser(1<<20, time.Second*2, false, nil) b.ResetTimer() b.ReportMetric(float64(1000), "packets/op") for i := 0; i < b.N; i++ { @@ -315,7 +297,7 @@ func BenchmarkMessageParserWithHint(b *testing.B) { packets[i] = GetPackets(false, 1, 1, buf[i])[0] } - parser := NewMessageParser(1<<30, time.Second*10, nil) + parser := NewMessageParser(1<<30, time.Second*10, false, nil) parser.Start = func(pckt *Packet) (bool, bool) { return false, proto.HasResponseTitle(pckt.Payload) } @@ -337,6 +319,6 @@ func BenchmarkNewAndParsePacket(b *testing.B) { data := append(generateHeader(true, 1024, 10), make([]byte, 10)...) b.ResetTimer() for i := 0; i < b.N; i++ { - ParsePacket(data, int(layers.LinkTypeLoop), 4, &gopacket.CaptureInfo{}) + ParsePacket(data, int(layers.LinkTypeLoop), 4, &gopacket.CaptureInfo{}, true) } }