use locks around client maps (#126)

Co-authored-by: Luca Moser <moser.luca@gmail.com>
This commit is contained in:
muXxer
2021-08-10 10:46:38 +08:00
committed by GitHub
co-authored by Luca Moser
parent c75ef2d6aa
commit 1d6979189a
2 changed files with 106 additions and 66 deletions
+18 -14
View File
@@ -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) {
+88 -52
View File
@@ -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 {