From 1d6979189a22ed8a775ee0213270dd3926137d28 Mon Sep 17 00:00:00 2001 From: muXxer Date: Tue, 10 Aug 2021 04:46:38 +0200 Subject: [PATCH] use locks around client maps (#126) Co-authored-by: Luca Moser --- broker/broker.go | 32 ++++++----- broker/client.go | 140 +++++++++++++++++++++++++++++------------------ 2 files changed, 106 insertions(+), 66 deletions(-) diff --git a/broker/broker.go b/broker/broker.go index a6d420c..66c373b 100644 --- a/broker/broker.go +++ b/broker/broker.go @@ -549,13 +549,19 @@ func (b *Broker) SendLocalSubsToRouter(c *client) { subInfo := packets.NewControlPacket(packets.Subscribe).(*packets.SubscribePacket) b.clients.Range(func(key, value interface{}) bool { client, ok := value.(*client) - if ok { - subs := client.subMap - for _, sub := range subs { - subInfo.Topics = append(subInfo.Topics, sub.topic) - subInfo.Qoss = append(subInfo.Qoss, sub.qos) - } + if !ok { + return true } + + client.subMapMu.RLock() + defer client.subMapMu.RUnlock() + + subs := client.subMap + for _, sub := range subs { + subInfo.Topics = append(subInfo.Topics, sub.topic) + subInfo.Qoss = append(subInfo.Qoss, sub.qos) + } + return true }) if len(subInfo.Topics) > 0 { @@ -629,16 +635,14 @@ func (b *Broker) PublishMessage(packet *packets.PublishPacket) { } } -func (b *Broker) BroadcastUnSubscribe(subs map[string]*subscription) { +func (b *Broker) BroadcastUnSubscribe(topicsToUnSubscribeFrom []string) { + if len(topicsToUnSubscribeFrom) == 0 { + return + } unsub := packets.NewControlPacket(packets.Unsubscribe).(*packets.UnsubscribePacket) - for topic, _ := range subs { - unsub.Topics = append(unsub.Topics, topic) - } - - if len(unsub.Topics) > 0 { - b.BroadcastSubOrUnsubMessage(unsub) - } + unsub.Topics = append(unsub.Topics, topicsToUnSubscribeFrom...) + b.BroadcastSubOrUnsubMessage(unsub) } func (b *Broker) OnlineOfflineNotification(clientID string, online bool) { diff --git a/broker/client.go b/broker/client.go index 955254e..335197c 100644 --- a/broker/client.go +++ b/broker/client.go @@ -66,12 +66,15 @@ type client struct { cancelFunc context.CancelFunc session *sessions.Session subMap map[string]*subscription + subMapMu sync.RWMutex topicsMgr *topics.Manager subs []interface{} qoss []byte rmsgs []*packets.PublishPacket routeSubMap map[string]uint64 + routeSubMapMu sync.Mutex awaitingRel map[uint16]int64 + awaitingRelMu sync.RWMutex maxAwaitingRel int inflight map[uint16]*inflightElem inflightMu sync.RWMutex @@ -273,7 +276,7 @@ func ProcessMessage(msg *Message) { // If a Server or Client receives a Control Packet // containing ill-formed UTF-8 it MUST close the Network Connection - c.conn.Close() + _ = c.conn.Close() // Update client status //c.status = Disconnected @@ -323,7 +326,7 @@ func ProcessMessage(msg *Message) { } case *packets.PubrelPacket: packet := ca.(*packets.PubrelPacket) - c.pubRel(packet.MessageID) + _ = c.pubRel(packet.MessageID) pubcomp := packets.NewControlPacket(packets.Pubcomp).(*packets.PubcompPacket) pubcomp.MessageID = packet.MessageID if err := c.WriterPacket(pubcomp); err != nil { @@ -525,14 +528,15 @@ func (c *client) processClientSubscribe(packet *packets.SubscribePacket) { if b == nil { return } - topics := packet.Topics + + subTopics := packet.Topics qoss := packet.Qoss suback := packets.NewControlPacket(packets.Suback).(*packets.SubackPacket) suback.MessageID = packet.MessageID var retcodes []byte - for i, topic := range topics { + for i, topic := range subTopics { t := topic //check topic auth for client if !b.CheckTopicAuth(SUB, c.info.clientID, c.info.username, c.info.remoteIP, topic) { @@ -562,10 +566,12 @@ func (c *client) processClientSubscribe(packet *packets.SubscribePacket) { topic = substr[2] } + c.subMapMu.Lock() if oldSub, exist := c.subMap[t]; exist { - c.topicsMgr.Unsubscribe([]byte(oldSub.topic), oldSub) + _ = c.topicsMgr.Unsubscribe([]byte(oldSub.topic), oldSub) delete(c.subMap, t) } + c.subMapMu.Unlock() sub := &subscription{ topic: topic, @@ -582,12 +588,13 @@ func (c *client) processClientSubscribe(packet *packets.SubscribePacket) { continue } + c.subMapMu.Lock() c.subMap[t] = sub + c.subMapMu.Unlock() - c.session.AddTopic(t, qoss[i]) + _ = c.session.AddTopic(t, qoss[i]) retcodes = append(retcodes, rqos) - c.topicsMgr.Retained([]byte(topic), &c.rmsgs) - + _ = c.topicsMgr.Retained([]byte(topic), &c.rmsgs) } suback.ReturnCodes = retcodes @@ -619,14 +626,15 @@ func (c *client) processRouterSubscribe(packet *packets.SubscribePacket) { if b == nil { return } - topics := packet.Topics + + subTopics := packet.Topics qoss := packet.Qoss suback := packets.NewControlPacket(packets.Suback).(*packets.SubackPacket) suback.MessageID = packet.MessageID var retcodes []byte - for i, topic := range topics { + for i, topic := range subTopics { t := topic groupName := "" share := false @@ -656,8 +664,13 @@ func (c *client) processRouterSubscribe(packet *packets.SubscribePacket) { continue } + c.subMapMu.Lock() c.subMap[t] = sub + c.subMapMu.Unlock() + + c.routeSubMapMu.Lock() addSubMap(c.routeSubMap, topic) + c.routeSubMapMu.Unlock() retcodes = append(retcodes, rqos) } @@ -687,20 +700,24 @@ func (c *client) processRouterUnSubscribe(packet *packets.UnsubscribePacket) { if b == nil { return } - topics := packet.Topics - for _, topic := range topics { - sub, exist := c.subMap[topic] - if exist { - retainNum := delSubMap(c.routeSubMap, topic) - if retainNum > 0 { + unSubTopics := packet.Topics + + for _, topic := range unSubTopics { + c.subMapMu.Lock() + if sub, exist := c.subMap[topic]; exist { + c.routeSubMapMu.Lock() + if retainNum := delSubMap(c.routeSubMap, topic); retainNum > 0 { + c.routeSubMapMu.Unlock() + c.subMapMu.Unlock() continue } + c.routeSubMapMu.Unlock() - c.topicsMgr.Unsubscribe([]byte(sub.topic), sub) + _ = c.topicsMgr.Unsubscribe([]byte(sub.topic), sub) delete(c.subMap, topic) } - + c.subMapMu.Unlock() } unsuback := packets.NewControlPacket(packets.Unsuback).(*packets.UnsubackPacket) @@ -721,9 +738,10 @@ func (c *client) processClientUnSubscribe(packet *packets.UnsubscribePacket) { if b == nil { return } - topics := packet.Topics - for _, topic := range topics { + unSubTopics := packet.Topics + + for _, topic := range unSubTopics { { //publish kafka @@ -737,12 +755,14 @@ func (c *client) processClientUnSubscribe(packet *packets.UnsubscribePacket) { } + c.subMapMu.Lock() sub, exist := c.subMap[topic] if exist { - c.topicsMgr.Unsubscribe([]byte(sub.topic), sub) - c.session.RemoveTopic(topic) + _ = c.topicsMgr.Unsubscribe([]byte(sub.topic), sub) + _ = c.session.RemoveTopic(topic) delete(c.subMap, topic) } + c.subMapMu.Unlock() } @@ -791,44 +811,52 @@ func (c *client) Close() { }) if c.conn != nil { - c.conn.Close() + _ = c.conn.Close() c.conn = nil } - subs := c.subMap + if b == nil { + return + } - if b != nil { - b.removeClient(c) - for _, sub := range subs { - // guard against race condition where a client gets Close() but wasn't initialized yet fully - if sub == nil || b.topicsMgr == nil { - continue - } - err := b.topicsMgr.Unsubscribe([]byte(sub.topic), sub) - if err != nil { - log.Error("unsubscribe error, ", zap.Error(err), zap.String("ClientID", c.info.clientID)) - } + b.removeClient(c) + + c.subMapMu.RLock() + defer c.subMapMu.RUnlock() + + unSubTopics := make([]string, 0) + for topic, sub := range c.subMap { + unSubTopics = append(unSubTopics, topic) + + // guard against race condition where a client gets Close() but wasn't initialized yet fully + if sub == nil || b.topicsMgr == nil { + continue } - if c.typ == CLIENT { - b.BroadcastUnSubscribe(subs) - //offline notification - b.OnlineOfflineNotification(c.info.clientID, false) - } - - if c.info.willMsg != nil { - b.PublishMessage(c.info.willMsg) - } - - if c.typ == CLUSTER { - b.ConnectToDiscovery() - } - - //do reconnect - if c.typ == REMOTE { - go b.connectRouter(c.route.remoteID, c.route.remoteUrl) + if err := b.topicsMgr.Unsubscribe([]byte(sub.topic), sub); err != nil { + log.Error("unsubscribe error, ", zap.Error(err), zap.String("ClientID", c.info.clientID)) } } + + if c.typ == CLIENT { + b.BroadcastUnSubscribe(unSubTopics) + //offline notification + b.OnlineOfflineNotification(c.info.clientID, false) + } + + if c.info.willMsg != nil { + b.PublishMessage(c.info.willMsg) + } + + if c.typ == CLUSTER { + b.ConnectToDiscovery() + } + + //do reconnect + if c.typ == REMOTE { + go b.connectRouter(c.route.remoteID, c.route.remoteUrl) + } + } func (c *client) WriterPacket(packet packets.ControlPacket) error { @@ -860,6 +888,8 @@ func (c *client) registerPublishPacketId(packetId uint16) error { return errors.New("DROPPED_QOS2_PACKET_FOR_TOO_MANY_AWAITING_REL") } + c.awaitingRelMu.Lock() + defer c.awaitingRelMu.Unlock() if _, found := c.awaitingRel[packetId]; found { return errors.New("RC_PACKET_IDENTIFIER_IN_USE") } @@ -869,6 +899,8 @@ func (c *client) registerPublishPacketId(packetId uint16) error { } func (c *client) isAwaitingFull() bool { + c.awaitingRelMu.RLock() + defer c.awaitingRelMu.RUnlock() if c.maxAwaitingRel == 0 { return false } @@ -879,6 +911,8 @@ func (c *client) isAwaitingFull() bool { } func (c *client) expireAwaitingRel() { + c.awaitingRelMu.Lock() + defer c.awaitingRelMu.Unlock() if len(c.awaitingRel) == 0 { return } @@ -892,6 +926,8 @@ func (c *client) expireAwaitingRel() { } func (c *client) pubRel(packetId uint16) error { + c.awaitingRelMu.Lock() + defer c.awaitingRelMu.Unlock() if _, found := c.awaitingRel[packetId]; found { delete(c.awaitingRel, packetId) } else {