From 85e84ce330fb54af7056abb79708eff8a0e88052 Mon Sep 17 00:00:00 2001 From: Leonid Bugaev Date: Wed, 11 May 2016 14:26:11 +0500 Subject: [PATCH] Improve handling of fragmented messages --- input_raw.go | 2 +- raw_socket_listener/listener.go | 48 +++++++++++++++---------- raw_socket_listener/listener_test.go | 12 +++++-- raw_socket_listener/tcp_message.go | 15 ++++---- raw_socket_listener/tcp_message_test.go | 12 +++---- 5 files changed, 55 insertions(+), 34 deletions(-) diff --git a/input_raw.go b/input_raw.go index e98d94c..10090bb 100644 --- a/input_raw.go +++ b/input_raw.go @@ -52,7 +52,7 @@ func (i *RAWInput) Read(data []byte) (int, error) { if msg.IsIncoming { header = payloadHeader(RequestPayload, msg.UUID(), msg.Start.UnixNano()) } else { - header = payloadHeader(ResponsePayload, msg.UUID(), msg.End.UnixNano()-msg.RequestStart.UnixNano()) + header = payloadHeader(ResponsePayload, msg.UUID(), msg.End.UnixNano()-msg.AssocMessage.Start.UnixNano()) } copy(data[0:len(header)], header) diff --git a/raw_socket_listener/listener.go b/raw_socket_listener/listener.go index b2bf26e..910e02c 100644 --- a/raw_socket_listener/listener.go +++ b/raw_socket_listener/listener.go @@ -45,7 +45,7 @@ type Listener struct { seqWithData map[uint32]uint32 // Ack -> Req - respAliases map[uint32]*request + respAliases map[uint32]*TCPMessage // Ack -> ID respWithoutReq map[uint32]tcpID @@ -90,7 +90,7 @@ func NewListener(addr string, port string, engine int, expire time.Duration) (l l.messages = make(map[tcpID]*TCPMessage) l.ackAliases = make(map[uint32]uint32) l.seqWithData = make(map[uint32]uint32) - l.respAliases = make(map[uint32]*request) + l.respAliases = make(map[uint32]*TCPMessage) l.respWithoutReq = make(map[uint32]tcpID) l.addr = addr @@ -167,31 +167,33 @@ func (t *Listener) dispatchMessage(message *TCPMessage) { t.deleteMessage(message) - log.Println("Dispatching, message", message.Seq, message.Ack, message.RequestAck, string(message.Bytes())) + // log.Println("Dispatching, message", message.Start.UnixNano(), message.Seq, message.Ack, string(message.Bytes())) if message.IsIncoming { // If there were response before request // log.Println("Looking for Response: ", t.respWithoutReq, message.ResponseAck) if respID, ok := t.respWithoutReq[message.ResponseAck]; ok { if resp, rok := t.messages[respID]; rok { - if resp.RequestAck == 0 { + // if resp.AssocMessage == nil { // log.Println("FOUND RESPONSE") - resp.RequestAck = message.Ack - resp.RequestStart = message.Start + resp.AssocMessage = message + message.AssocMessage = resp if resp.IsFinished() { defer t.dispatchMessage(resp) } - } + // } } } - + if resp, ok := t.messages[message.ResponseID]; ok { + resp.AssocMessage = message + } } else { - if message.RequestAck == 0 { + if message.AssocMessage == nil { if responseRequest, ok := t.respAliases[message.Ack]; ok { - message.RequestStart = responseRequest.start - message.RequestAck = responseRequest.ack + message.AssocMessage = responseRequest + responseRequest.AssocMessage = message } } @@ -199,7 +201,7 @@ func (t *Listener) dispatchMessage(message *TCPMessage) { delete(t.respWithoutReq, message.Ack) // Do not track responses which have no associated requests - if message.RequestAck == 0 { + if message.AssocMessage == nil { // log.Println("Can't dispatch resp", message.Seq, message.Ack, string(message.Bytes())) return } @@ -399,6 +401,10 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { if m.Ack == packet.Ack && bytes.Equal(m.packets[0].Addr, packet.Addr) { t.deleteMessage(m) + if m.AssocMessage != nil { + m.AssocMessage.AssocMessage = nil + } + for _, pkt := range m.packets { // log.Println("Updating ack", parentAck, pkt.Ack) pkt.UpdateAck(parentAck) @@ -416,7 +422,7 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { packet.UpdateAck(alias) } - var responseRequest *request + var responseRequest *TCPMessage if !isIncoming { responseRequest, _ = t.respAliases[packet.Ack] @@ -430,9 +436,8 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { if !isIncoming { if responseRequest != nil { - message.RequestStart = responseRequest.start - message.RequestAck = responseRequest.ack - message.RequestID = responseRequest.id + message.AssocMessage = responseRequest + responseRequest.AssocMessage = message } else { t.respWithoutReq[packet.Ack] = packet.ID } @@ -454,6 +459,9 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { for _, m := range t.messages { if m.Seq == seq { t.deleteMessage(m) + if m.AssocMessage != nil { + message.AssocMessage = m.AssocMessage + } // log.Println("2: Adding ack alias:", m.Ack, packet.Ack) t.ackAliases[m.Ack] = packet.Ack @@ -480,7 +488,7 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { } message.UpdateResponseAck() - t.respAliases[message.ResponseAck] = &request{message.ID(), message.Start, message.Ack} + t.respAliases[message.ResponseAck] = message } // If message contains only single packet immediately dispatch it @@ -493,7 +501,11 @@ func (t *Listener) processTCPPacket(packet *TCPPacket) { } } } else { - if req, ok := t.messages[message.RequestID]; ok { + if message.AssocMessage == nil { + return + } + + if req, ok := t.messages[message.AssocMessage.ID()]; ok { if req.IsFinished() { t.dispatchMessage(req) t.dispatchMessage(message) diff --git a/raw_socket_listener/listener_test.go b/raw_socket_listener/listener_test.go index ff6a6f7..5e7f06c 100644 --- a/raw_socket_listener/listener_test.go +++ b/raw_socket_listener/listener_test.go @@ -257,6 +257,14 @@ func testChunkedSequence(t *testing.T, listener *Listener, packets ...*TCPPacket time.Sleep(20 * time.Millisecond) + if len(listener.packetsChan) != 0 { + t.Fatal("packetsChan non empty:", listener.packetsChan) + } + + if len(listener.messagesChan) != 0 { + t.Fatal("messagesChan non empty:", <- listener.messagesChan) + } + if len(listener.messages) != 0 { t.Fatal("Messages non empty:", listener.messages) } @@ -312,13 +320,13 @@ func TestRawListenerChunkedWrongOrder(t *testing.T) { // Should re-construct message from all possible combinations for i := 0; i < 6*5*4*3*2*1; i++ { - if i != 87 { + if i < 54 || i > 57 { continue } packets := permutation(i, []*TCPPacket{reqPacket1, reqPacket2, reqPacket3, reqPacket4, respPacket1, respPacket2}) - t.Log("permutation:", i, packets) + t.Log("permutation:", i) testChunkedSequence(t, listener, packets...) } } diff --git a/raw_socket_listener/tcp_message.go b/raw_socket_listener/tcp_message.go index aea2b54..2e39984 100644 --- a/raw_socket_listener/tcp_message.go +++ b/raw_socket_listener/tcp_message.go @@ -23,12 +23,11 @@ type TCPMessage struct { Seq uint32 Ack uint32 ResponseAck uint32 - RequestStart time.Time + ResponseID tcpID DataAck uint32 DataSeq uint32 - RequestAck uint32 - RequestID tcpID - ResponseID tcpID + + AssocMessage *TCPMessage Start time.Time End time.Time IsIncoming bool @@ -153,7 +152,7 @@ func (t *TCPMessage) IsFinished() bool { } else { // Request not found // Can be because response came first or request request was just missing - if t.RequestAck == 0 { + if t.AssocMessage == nil { return false } @@ -182,11 +181,13 @@ func (t *TCPMessage) UUID() []byte { var key []byte if t.IsIncoming { + // log.Println("UUID:", t.Ack, t.Start.UnixNano()) key = strconv.AppendInt(key, t.Start.UnixNano(), 10) key = strconv.AppendUint(key, uint64(t.Ack), 10) } else { - key = strconv.AppendInt(key, t.RequestStart.UnixNano(), 10) - key = strconv.AppendUint(key, uint64(t.RequestAck), 10) + // log.Println("RequestMessage:", t.AssocMessage.Ack, t.AssocMessage.Start.UnixNano()) + key = strconv.AppendInt(key, t.AssocMessage.Start.UnixNano(), 10) + key = strconv.AppendUint(key, uint64(t.AssocMessage.Ack), 10) } uuid := make([]byte, 40) diff --git a/raw_socket_listener/tcp_message_test.go b/raw_socket_listener/tcp_message_test.go index 9bc4a50..2df079e 100644 --- a/raw_socket_listener/tcp_message_test.go +++ b/raw_socket_listener/tcp_message_test.go @@ -114,40 +114,40 @@ func TestTCPMessageIsFinished(t *testing.T) { // Responses msg = buildMessage(buildPacket(false, 1, 1, []byte("HTTP/1.1 200 OK\r\n\r\n"))) - msg.RequestAck = 1 + msg.AssocMessage = &TCPMessage{} if !msg.IsFinished() { t.Error("Should mark simple response as finished") } msg = buildMessage(buildPacket(false, 1, 1, []byte("HTTP/1.1 200 OK\r\n\r\n"))) - msg.RequestAck = 0 + msg.AssocMessage = nil if msg.IsFinished() { t.Error("Should not mark responses without associated requests") } msg = buildMessage(buildPacket(false, 1, 1, []byte("HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n"))) - msg.RequestAck = 1 + msg.AssocMessage = &TCPMessage{} if msg.IsFinished() { t.Error("Should mark chunked response as non finished") } msg = buildMessage(buildPacket(false, 1, 1, []byte("HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n"))) - msg.RequestAck = 1 + msg.AssocMessage = &TCPMessage{} if !msg.IsFinished() { t.Error("Should mark Content-Length: 0 respones as finished") } msg = buildMessage(buildPacket(false, 1, 1, []byte("HTTP/1.1 200 OK\r\nContent-Length: 1\r\n\r\na"))) - msg.RequestAck = 1 + msg.AssocMessage = &TCPMessage{} if !msg.IsFinished() { t.Error("Should mark valid Content-Length respones as finished") } msg = buildMessage(buildPacket(false, 1, 1, []byte("HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\na"))) - msg.RequestAck = 1 + msg.AssocMessage = &TCPMessage{} if msg.IsFinished() { t.Error("Should not mark not valid Content-Length respones as finished")