mirror of
https://github.com/buger/goreplay.git
synced 2024-04-21 12:32:02 +00:00
Refactor concurrency to use more channels
This commit is contained in:
+12
-14
@@ -44,25 +44,23 @@ func Run() {
|
||||
listener := RAWTCPListen(Settings.address, Settings.port)
|
||||
|
||||
for {
|
||||
message := listener.Receive()
|
||||
m := listener.Receive()
|
||||
|
||||
go func(m *TCPMessage) {
|
||||
if Settings.verbose {
|
||||
buf := bytes.NewBuffer(m.Bytes())
|
||||
reader := bufio.NewReader(buf)
|
||||
if Settings.verbose {
|
||||
buf := bytes.NewBuffer(m.Bytes())
|
||||
reader := bufio.NewReader(buf)
|
||||
|
||||
request, err := http.ReadRequest(reader)
|
||||
request, err := http.ReadRequest(reader)
|
||||
|
||||
if err != nil {
|
||||
Debug("Error while parsing request:", string(m.Bytes()))
|
||||
} else {
|
||||
request.ParseMultipartForm(32 << 20)
|
||||
Debug("Forwarding request:", request)
|
||||
}
|
||||
if err != nil {
|
||||
Debug("Error while parsing request:", string(m.Bytes()))
|
||||
} else {
|
||||
request.ParseMultipartForm(32 << 20)
|
||||
Debug("Forwarding request:", request)
|
||||
}
|
||||
}
|
||||
|
||||
conn.Write(m.Bytes())
|
||||
}(message)
|
||||
conn.Write(m.Bytes())
|
||||
}
|
||||
|
||||
conn.Close()
|
||||
|
||||
@@ -4,25 +4,28 @@ import (
|
||||
"encoding/binary"
|
||||
"log"
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
type RAWTCPListener struct {
|
||||
messages map[uint32]*TCPMessage // buffer of TCPMessages waiting to be send
|
||||
messages []*TCPMessage // buffer of TCPMessages waiting to be send
|
||||
|
||||
c_packets chan *TCPPacket
|
||||
c_messages chan *TCPMessage
|
||||
|
||||
c_add_message chan *TCPMessage
|
||||
c_del_message chan *TCPMessage
|
||||
|
||||
addr string
|
||||
port int
|
||||
}
|
||||
|
||||
func RAWTCPListen(addr string, port int) (listener *RAWTCPListener) {
|
||||
listener = &RAWTCPListener{}
|
||||
listener.messages = make(map[uint32]*TCPMessage)
|
||||
|
||||
listener.c_packets = make(chan *TCPPacket)
|
||||
listener.c_messages = make(chan *TCPMessage)
|
||||
listener.c_add_message = make(chan *TCPMessage)
|
||||
listener.c_del_message = make(chan *TCPMessage)
|
||||
|
||||
listener.addr = addr
|
||||
listener.port = port
|
||||
@@ -34,32 +37,41 @@ func RAWTCPListen(addr string, port int) (listener *RAWTCPListener) {
|
||||
}
|
||||
|
||||
func (t *RAWTCPListener) listen() {
|
||||
|
||||
for {
|
||||
var messages chan *TCPMessage
|
||||
var message *TCPMessage
|
||||
|
||||
for _, msg := range t.messages {
|
||||
if msg.Complete() {
|
||||
messages = t.c_messages
|
||||
message = msg
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
select {
|
||||
case messages <- message:
|
||||
delete(t.messages, message.ack)
|
||||
case message := <-t.c_del_message:
|
||||
t.deleteMessage(message)
|
||||
Debug("Deleted")
|
||||
t.c_messages <- message
|
||||
|
||||
case packet := <-t.c_packets:
|
||||
t.processTCPPacket(packet)
|
||||
|
||||
// Ensure that this will be run at least each 200 ms, to ensure that all messages will be send
|
||||
// Without it last message may not be send (it will be send only on next TCP packets)
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
Debug("Processed")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *RAWTCPListener) deleteMessage(message *TCPMessage) bool {
|
||||
var idx int = -1
|
||||
|
||||
for i, m := range t.messages {
|
||||
if m.Ack == message.Ack {
|
||||
idx = i
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if idx == -1 {
|
||||
return false
|
||||
}
|
||||
|
||||
copy(t.messages[idx:], t.messages[idx+1:])
|
||||
t.messages[len(t.messages)-1] = nil // or the zero value of T
|
||||
t.messages = t.messages[:len(t.messages)-1]
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *RAWTCPListener) readTCPPackets() {
|
||||
conn, e := net.ListenPacket("ip4:tcp", t.addr)
|
||||
defer conn.Close()
|
||||
@@ -93,7 +105,6 @@ func (t *RAWTCPListener) readTCPPackets() {
|
||||
copy(new_buf, buf[:n])
|
||||
|
||||
packet := NewTCPPacket(new_buf)
|
||||
|
||||
t.c_packets <- packet
|
||||
}
|
||||
}
|
||||
@@ -103,13 +114,23 @@ func (t *RAWTCPListener) readTCPPackets() {
|
||||
|
||||
//
|
||||
func (t *RAWTCPListener) processTCPPacket(packet *TCPPacket) {
|
||||
ack := packet.Ack
|
||||
var message *TCPMessage
|
||||
|
||||
if _, ok := t.messages[ack]; !ok {
|
||||
t.messages[ack] = NewTCPMessage(ack)
|
||||
for _, msg := range t.messages {
|
||||
if msg.Ack == packet.Ack {
|
||||
message = msg
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
t.messages[ack].AddPacket(packet)
|
||||
if message == nil {
|
||||
message = NewTCPMessage(packet.Ack, t.c_del_message)
|
||||
Debug("Adding message")
|
||||
|
||||
t.messages = append(t.messages, message)
|
||||
}
|
||||
|
||||
message.c_packets <- packet
|
||||
}
|
||||
|
||||
func (t *RAWTCPListener) Receive() *TCPMessage {
|
||||
|
||||
+59
-20
@@ -5,6 +5,8 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
const MSG_EXPIRE = 200 * time.Millisecond
|
||||
|
||||
// TCPMessage ensure that all TCP packets for given request is received, and processed in right sequence
|
||||
// Its needed because all TCP message can be fragmented or re-transmitted
|
||||
//
|
||||
@@ -12,19 +14,50 @@ import (
|
||||
// Message can be compiled from unique packets with same message_id which sorted by sequence
|
||||
// Message is received if we did't receive any packets for 200ms OR if we received packet with "fin" flag
|
||||
type TCPMessage struct {
|
||||
ack uint32 // Message ID
|
||||
packets map[int]*TCPPacket // map[packet.sequence]*TCPPacket
|
||||
updated int64 // time of last packet
|
||||
Ack uint32 // Message ID
|
||||
packets []*TCPPacket
|
||||
|
||||
timer *time.Timer
|
||||
|
||||
expired bool
|
||||
|
||||
c_packets chan *TCPPacket
|
||||
c_closing chan int
|
||||
|
||||
c_del_message chan *TCPMessage
|
||||
}
|
||||
|
||||
func NewTCPMessage(ack uint32) (msg *TCPMessage) {
|
||||
msg = &TCPMessage{}
|
||||
msg.packets = make(map[int]*TCPPacket)
|
||||
msg.updated = time.Now().UnixNano()
|
||||
msg.ack = ack
|
||||
func NewTCPMessage(Ack uint32, c_del chan *TCPMessage) (msg *TCPMessage) {
|
||||
msg = &TCPMessage{Ack: Ack}
|
||||
|
||||
msg.c_packets = make(chan *TCPPacket)
|
||||
msg.c_closing = make(chan int)
|
||||
msg.c_del_message = c_del
|
||||
|
||||
msg.timer = time.AfterFunc(MSG_EXPIRE, msg.Timeout)
|
||||
|
||||
go msg.ListenPackets()
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func (t *TCPMessage) ListenPackets() {
|
||||
for {
|
||||
select {
|
||||
case <-t.c_closing:
|
||||
close(t.c_packets)
|
||||
return
|
||||
case packet := <-t.c_packets:
|
||||
t.AddPacket(packet)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *TCPMessage) Timeout() {
|
||||
t.c_closing <- 1
|
||||
t.c_del_message <- t
|
||||
}
|
||||
|
||||
// Sort packets in right orders and return message content
|
||||
func (t *TCPMessage) Bytes() (output []byte) {
|
||||
mk := make([]int, len(t.packets))
|
||||
@@ -46,19 +79,25 @@ func (t *TCPMessage) Bytes() (output []byte) {
|
||||
|
||||
// Add packet to the message
|
||||
func (t *TCPMessage) AddPacket(packet *TCPPacket) {
|
||||
seq := int(packet.Seq)
|
||||
|
||||
if _, ok := t.packets[seq]; !ok {
|
||||
t.packets[seq] = packet
|
||||
} else {
|
||||
Debug("Received packet with same sequence")
|
||||
if t.expired {
|
||||
Debug("Adding packet to expired message")
|
||||
return
|
||||
}
|
||||
|
||||
t.updated = time.Now().UnixNano()
|
||||
}
|
||||
packetFound := false
|
||||
|
||||
// TCP message is complete if we not received any packets for 200ms since last packet
|
||||
func (t *TCPMessage) Complete() bool {
|
||||
ns := time.Now().UnixNano()
|
||||
return (ns - t.updated) > int64(200*time.Millisecond)
|
||||
for _, pkt := range t.packets {
|
||||
if packet.Seq == pkt.Seq {
|
||||
packetFound = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if packetFound {
|
||||
Debug("Received packet with same sequence")
|
||||
} else {
|
||||
t.packets = append(t.packets, packet)
|
||||
}
|
||||
|
||||
t.timer.Reset(MSG_EXPIRE)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user