From eebff88a29a8412330d446ed23db69b990b29224 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Mon, 28 Jul 2025 16:57:58 +0800 Subject: [PATCH 01/19] fix: performance issues with Kafka caused by encapsulating the MQ interface --- go.mod | 2 +- go.sum | 4 ++-- internal/msgtransfer/init.go | 6 +++--- internal/msgtransfer/online_history_msg_handler.go | 13 ++++++++++--- internal/msgtransfer/online_msg_to_mongo_handler.go | 10 +++++++--- internal/push/push.go | 9 +++++---- 6 files changed, 28 insertions(+), 16 deletions(-) diff --git a/go.mod b/go.mod index c06451aaa..775765706 100644 --- a/go.mod +++ b/go.mod @@ -13,7 +13,7 @@ require ( github.com/grpc-ecosystem/go-grpc-prometheus v1.2.0 github.com/mitchellh/mapstructure v1.5.0 github.com/openimsdk/protocol v0.0.73-alpha.12 - github.com/openimsdk/tools v0.0.50-alpha.96 + github.com/openimsdk/tools v0.0.50-alpha.97 github.com/pkg/errors v0.9.1 // indirect github.com/prometheus/client_golang v1.18.0 github.com/stretchr/testify v1.9.0 diff --git a/go.sum b/go.sum index d29eb0f3f..329a916ec 100644 --- a/go.sum +++ b/go.sum @@ -349,8 +349,8 @@ github.com/openimsdk/gomake v0.0.15-alpha.11 h1:PQudYDRESYeYlUYrrLLJhYIlUPO5x7FA github.com/openimsdk/gomake v0.0.15-alpha.11/go.mod h1:PndCozNc2IsQIciyn9mvEblYWZwJmAI+06z94EY+csI= github.com/openimsdk/protocol v0.0.73-alpha.12 h1:2NYawXeHChYUeSme6QJ9pOLh+Empce2WmwEtbP4JvKk= github.com/openimsdk/protocol v0.0.73-alpha.12/go.mod h1:WF7EuE55vQvpyUAzDXcqg+B+446xQyEba0X35lTINmw= -github.com/openimsdk/tools v0.0.50-alpha.96 h1:U44Fq2jHiEvGi9zuYAnTRNx3Xd9T7P/kBAZLHvQ8xg4= -github.com/openimsdk/tools v0.0.50-alpha.96/go.mod h1:n2poR3asX1e1XZce4O+MOWAp+X02QJRFvhcLCXZdzRo= +github.com/openimsdk/tools v0.0.50-alpha.97 h1:6ik5w3PpgDG6VjSo3nb3FT/fxN3JX7iIARVxVu9g7VY= +github.com/openimsdk/tools v0.0.50-alpha.97/go.mod h1:n2poR3asX1e1XZce4O+MOWAp+X02QJRFvhcLCXZdzRo= github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6Wq+LM= github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs= github.com/pierrec/lz4/v4 v4.1.21 h1:yOVMLb6qSIDP67pl/5F7RepeKYu/VmTyEXvuMI5d9mQ= diff --git a/internal/msgtransfer/init.go b/internal/msgtransfer/init.go index bbec3f9a2..35026c79a 100644 --- a/internal/msgtransfer/init.go +++ b/internal/msgtransfer/init.go @@ -134,7 +134,7 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - historyMongoHandler := NewOnlineHistoryMongoConsumerHandler(msgTransferDatabase,config) + historyMongoHandler := NewOnlineHistoryMongoConsumerHandler(msgTransferDatabase, config) msgTransfer := &MsgTransfer{ historyConsumer: historyConsumer, @@ -161,8 +161,8 @@ func (m *MsgTransfer) Start(ctx context.Context) error { }() go func() { - fn := func(ctx context.Context, key string, value []byte) error { - m.historyMongoHandler.HandleChatWs2Mongo(ctx, key, value) + fn := func(msg mq.Message) error { + m.historyMongoHandler.HandleChatWs2Mongo(msg) return nil } for { diff --git a/internal/msgtransfer/online_history_msg_handler.go b/internal/msgtransfer/online_history_msg_handler.go index 05775a1e6..8b212774a 100644 --- a/internal/msgtransfer/online_history_msg_handler.go +++ b/internal/msgtransfer/online_history_msg_handler.go @@ -18,6 +18,7 @@ import ( "context" "encoding/json" "errors" + "github.com/openimsdk/tools/mq" "sync" "time" @@ -77,6 +78,7 @@ type ConsumerMessage struct { Ctx context.Context Key string Value []byte + Raw mq.Message } func NewOnlineHistoryRedisConsumerHandler(ctx context.Context, client discovery.Conn, config *Config, database controller.MsgTransferDatabase) (*OnlineHistoryRedisConsumerHandler, error) { @@ -113,6 +115,11 @@ func NewOnlineHistoryRedisConsumerHandler(ctx context.Context, client discovery. b.Do = och.do och.redisMessageBatches = b + och.redisMessageBatches.OnComplete = func(lastMessage *ConsumerMessage, totalCount int) { + lastMessage.Raw.Mark() + lastMessage.Raw.Commit() + } + return &och, nil } func (och *OnlineHistoryRedisConsumerHandler) do(ctx context.Context, channelID int, val *batcher.Msg[ConsumerMessage]) { @@ -388,10 +395,10 @@ func withAggregationCtx(ctx context.Context, values []*ContextMsg) context.Conte return mcontext.SetOperationID(ctx, allMessageOperationID) } -func (och *OnlineHistoryRedisConsumerHandler) HandlerRedisMessage(ctx context.Context, key string, value []byte) error { // a instance in the consumer group - err := och.redisMessageBatches.Put(ctx, &ConsumerMessage{Ctx: ctx, Key: key, Value: value}) +func (och *OnlineHistoryRedisConsumerHandler) HandlerRedisMessage(msg mq.Message) error { // a instance in the consumer group + err := och.redisMessageBatches.Put(msg.Context(), &ConsumerMessage{Ctx: msg.Context(), Key: msg.Key(), Value: msg.Value(), Raw: msg}) if err != nil { - log.ZWarn(ctx, "put msg to error", err, "key", key, "value", value) + log.ZWarn(msg.Context(), "put msg to error", err, "key", msg.Key(), "value", msg.Value()) } return nil } diff --git a/internal/msgtransfer/online_msg_to_mongo_handler.go b/internal/msgtransfer/online_msg_to_mongo_handler.go index 6c1498f82..8611af7ea 100644 --- a/internal/msgtransfer/online_msg_to_mongo_handler.go +++ b/internal/msgtransfer/online_msg_to_mongo_handler.go @@ -15,12 +15,12 @@ package msgtransfer import ( - "context" + "github.com/openimsdk/protocol/constant" + "github.com/openimsdk/tools/mq" "github.com/openimsdk/open-im-server/v3/pkg/common/prommetrics" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/controller" "github.com/openimsdk/open-im-server/v3/pkg/common/webhook" - "github.com/openimsdk/protocol/constant" pbmsg "github.com/openimsdk/protocol/msg" "github.com/openimsdk/tools/log" "google.golang.org/protobuf/proto" @@ -40,7 +40,10 @@ func NewOnlineHistoryMongoConsumerHandler(database controller.MsgTransferDatabas } } -func (mc *OnlineHistoryMongoConsumerHandler) HandleChatWs2Mongo(ctx context.Context, key string, msg []byte) { +func (mc *OnlineHistoryMongoConsumerHandler) HandleChatWs2Mongo(val mq.Message) { + ctx := val.Context() + key := val.Key() + msg := val.Value() msgFromMQ := pbmsg.MsgDataToMongoByMQ{} err := proto.Unmarshal(msg, &msgFromMQ) if err != nil { @@ -58,6 +61,7 @@ func (mc *OnlineHistoryMongoConsumerHandler) HandleChatWs2Mongo(ctx context.Cont prommetrics.MsgInsertMongoFailedCounter.Inc() } else { prommetrics.MsgInsertMongoSuccessCounter.Inc() + val.Mark() } for _, msgData := range msgFromMQ.MsgData { diff --git a/internal/push/push.go b/internal/push/push.go index 1d6f8cb30..bf95b6acc 100644 --- a/internal/push/push.go +++ b/internal/push/push.go @@ -2,6 +2,7 @@ package push import ( "context" + "github.com/openimsdk/tools/mq" "math/rand" "strconv" @@ -106,8 +107,8 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg go func() { pushHandler.WaitCache() - fn := func(ctx context.Context, key string, value []byte) error { - pushHandler.HandleMs2PsChat(authverify.WithTempAdmin(ctx), value) + fn := func(msg mq.Message) error { + pushHandler.HandleMs2PsChat(authverify.WithTempAdmin(msg.Context()), msg.Value()) return nil } consumerCtx := mcontext.SetOperationID(context.Background(), "push_"+strconv.Itoa(int(rand.Uint32()))) @@ -121,8 +122,8 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg }() go func() { - fn := func(ctx context.Context, key string, value []byte) error { - offlineHandler.HandleMsg2OfflinePush(ctx, value) + fn := func(msg mq.Message) error { + offlineHandler.HandleMsg2OfflinePush(msg.Context(), msg.Value()) return nil } consumerCtx := mcontext.SetOperationID(context.Background(), "push_"+strconv.Itoa(int(rand.Uint32()))) From d9c3504afd6bb4f8a93154215995d8acbde01d7f Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 1 Aug 2025 10:13:42 +0800 Subject: [PATCH 02/19] fix: admin token in standalone mode --- internal/api/router.go | 5 +++++ pkg/common/storage/cache/redis/token.go | 9 ++++----- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/internal/api/router.go b/internal/api/router.go index 8a4199581..1d3a92dd7 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -97,6 +97,11 @@ func newGinRouter(ctx context.Context, client discovery.SvcDiscoveryRegistry, cf case BestSpeed: r.Use(gzip.Gzip(gzip.BestSpeed)) } + if config.Standalone() { + r.Use(func(c *gin.Context) { + c.Set(authverify.CtxAdminUserIDsKey, cfg.Share.IMAdminUser.UserIDs) + }) + } r.Use(api.GinLogger(), prommetricsGin(), gin.RecoveryWithWriter(gin.DefaultErrorWriter, mw.GinPanicErr), mw.CorsHandler(), mw.GinParseOperationID(), GinParseToken(rpcli.NewAuthClient(authConn)), setGinIsAdmin(cfg.Share.IMAdminUser.UserIDs)) diff --git a/pkg/common/storage/cache/redis/token.go b/pkg/common/storage/cache/redis/token.go index b3870daee..c74ccce66 100644 --- a/pkg/common/storage/cache/redis/token.go +++ b/pkg/common/storage/cache/redis/token.go @@ -165,16 +165,15 @@ func (c *tokenCache) DeleteTokenByTokenMap(ctx context.Context, userID string, t } func (c *tokenCache) DeleteAndSetTemporary(ctx context.Context, userID string, platformID int, fields []string) error { - key := cachekey.GetTokenKey(userID, platformID) - if err := c.rdb.HDel(ctx, key, fields...).Err(); err != nil { - return errs.Wrap(err) - } for _, f := range fields { k := cachekey.GetTemporaryTokenKey(userID, platformID, f) if err := c.rdb.Set(ctx, k, "", time.Minute*5).Err(); err != nil { return errs.Wrap(err) } } - + key := cachekey.GetTokenKey(userID, platformID) + if err := c.rdb.HDel(ctx, key, fields...).Err(); err != nil { + return errs.Wrap(err) + } return nil } From 9a1d2a85cdb555d7eef4e79d2c866dacdfbd86ae Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Wed, 15 Oct 2025 10:10:29 +0800 Subject: [PATCH 03/19] fix: full id version --- internal/rpc/conversation/sync.go | 2 +- internal/rpc/group/sync.go | 4 ++-- internal/rpc/relation/sync.go | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/internal/rpc/conversation/sync.go b/internal/rpc/conversation/sync.go index a24dd85c6..85128f719 100644 --- a/internal/rpc/conversation/sync.go +++ b/internal/rpc/conversation/sync.go @@ -27,7 +27,7 @@ func (c *conversationServer) GetFullOwnerConversationIDs(ctx context.Context, re conversationIDs = nil } return &conversation.GetFullOwnerConversationIDsResp{ - Version: idHash, + Version: uint64(vl.Version), VersionID: vl.ID.Hex(), Equal: req.IdHash == idHash, ConversationIDs: conversationIDs, diff --git a/internal/rpc/group/sync.go b/internal/rpc/group/sync.go index b864fbf53..92c7a60ce 100644 --- a/internal/rpc/group/sync.go +++ b/internal/rpc/group/sync.go @@ -34,7 +34,7 @@ func (g *groupServer) GetFullGroupMemberUserIDs(ctx context.Context, req *pbgrou userIDs = nil } return &pbgroup.GetFullGroupMemberUserIDsResp{ - Version: idHash, + Version: uint64(vl.Version), VersionID: vl.ID.Hex(), Equal: req.IdHash == idHash, UserIDs: userIDs, @@ -58,7 +58,7 @@ func (g *groupServer) GetFullJoinGroupIDs(ctx context.Context, req *pbgroup.GetF groupIDs = nil } return &pbgroup.GetFullJoinGroupIDsResp{ - Version: idHash, + Version: uint64(vl.Version), VersionID: vl.ID.Hex(), Equal: req.IdHash == idHash, GroupIDs: groupIDs, diff --git a/internal/rpc/relation/sync.go b/internal/rpc/relation/sync.go index 79fa0858c..187f6238d 100644 --- a/internal/rpc/relation/sync.go +++ b/internal/rpc/relation/sync.go @@ -56,7 +56,7 @@ func (s *friendServer) GetFullFriendUserIDs(ctx context.Context, req *relation.G userIDs = nil } return &relation.GetFullFriendUserIDsResp{ - Version: idHash, + Version: uint64(vl.Version), VersionID: vl.ID.Hex(), Equal: req.IdHash == idHash, UserIDs: userIDs, From ebda95fb11abdddd52ee645dcb2d0697ac9b8e35 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 12 Dec 2025 15:23:45 +0800 Subject: [PATCH 04/19] fix: resolve deadlock in cache eviction and improve GetBatch implementation --- pkg/localcache/cache.go | 27 ++++++----- pkg/localcache/cache_test.go | 67 ++++++++++++++++++++++++++++ pkg/localcache/init.go | 4 -- pkg/localcache/lru/lru_expiration.go | 49 +++++++++++++++++++- pkg/localcache/lru/lru_slot.go | 6 +-- 5 files changed, 133 insertions(+), 20 deletions(-) diff --git a/pkg/localcache/cache.go b/pkg/localcache/cache.go index 07d36cf46..b2376d6f1 100644 --- a/pkg/localcache/cache.go +++ b/pkg/localcache/cache.go @@ -47,15 +47,15 @@ func New[V any](opts ...Option) Cache[V] { if opt.localSlotNum > 0 && opt.localSlotSize > 0 { createSimpleLRU := func() lru.LRU[string, V] { if opt.expirationEvict { - return lru.NewExpirationLRU(opt.localSlotSize, opt.localSuccessTTL, opt.localFailedTTL, opt.target, c.onEvict) + return lru.NewExpirationLRU[string, V](opt.localSlotSize, opt.localSuccessTTL, opt.localFailedTTL, opt.target, c.onEvict) } else { - return lru.NewLazyLRU(opt.localSlotSize, opt.localSuccessTTL, opt.localFailedTTL, opt.target, c.onEvict) + return lru.NewLazyLRU[string, V](opt.localSlotSize, opt.localSuccessTTL, opt.localFailedTTL, opt.target, c.onEvict) } } if opt.localSlotNum == 1 { c.local = createSimpleLRU() } else { - c.local = lru.NewSlotLRU(opt.localSlotNum, LRUStringHash, createSimpleLRU) + c.local = lru.NewSlotLRU[string, V](opt.localSlotNum, LRUStringHash, createSimpleLRU) } if opt.linkSlotNum > 0 { c.link = link.New(opt.linkSlotNum) @@ -71,14 +71,19 @@ type cache[V any] struct { } func (c *cache[V]) onEvict(key string, value V) { - _ = value - if c.link != nil { - lks := c.link.Del(key) - for k := range lks { - if key != k { // prevent deadlock - c.local.Del(k) - } + // Do not delete other keys while the underlying LRU still holds its lock; + // defer linked deletions to avoid re-entering the same slot and deadlocking. + if lks := c.link.Del(key); len(lks) > 0 { + go c.delLinked(key, lks) + } + } +} + +func (c *cache[V]) delLinked(src string, keys map[string]struct{}) { + for k := range keys { + if src != k { + c.local.Del(k) } } } @@ -105,7 +110,7 @@ func (c *cache[V]) Get(ctx context.Context, key string, fetch func(ctx context.C func (c *cache[V]) GetLink(ctx context.Context, key string, fetch func(ctx context.Context) (V, error), link ...string) (V, error) { if c.local != nil { return c.local.Get(key, func() (V, error) { - if len(link) > 0 { + if len(link) > 0 && c.link != nil { c.link.Link(key, link...) } return fetch(ctx) diff --git a/pkg/localcache/cache_test.go b/pkg/localcache/cache_test.go index c206e6799..13eb20797 100644 --- a/pkg/localcache/cache_test.go +++ b/pkg/localcache/cache_test.go @@ -22,6 +22,8 @@ import ( "sync/atomic" "testing" "time" + + "github.com/openimsdk/open-im-server/v3/pkg/localcache/lru" ) func TestName(t *testing.T) { @@ -91,3 +93,68 @@ func TestName(t *testing.T) { t.Log("del", del.Load()) // 137.35s } + +// Test deadlock scenario when eviction callback deletes a linked key that hashes to the same slot. +func TestCacheEvictDeadlock(t *testing.T) { + ctx := context.Background() + c := New[string](WithLocalSlotNum(1), WithLocalSlotSize(1), WithLazy()) + + if _, err := c.GetLink(ctx, "k1", func(ctx context.Context) (string, error) { + return "v1", nil + }, "k2"); err != nil { + t.Fatalf("seed cache failed: %v", err) + } + + done := make(chan struct{}) + go func() { + defer close(done) + _, _ = c.GetLink(ctx, "k2", func(ctx context.Context) (string, error) { + return "v2", nil + }, "k1") + }() + + select { + case <-done: + // expected to finish quickly; current implementation deadlocks here. + case <-time.After(time.Second): + t.Fatal("GetLink deadlocked during eviction of linked key") + } +} + +func TestExpirationLRUGetBatch(t *testing.T) { + l := lru.NewExpirationLRU[string, string](2, time.Minute, time.Second*5, EmptyTarget{}, nil) + + keys := []string{"a", "b"} + values, err := l.GetBatch(keys, func(keys []string) (map[string]string, error) { + res := make(map[string]string) + for _, k := range keys { + res[k] = k + "_v" + } + return res, nil + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(values) != len(keys) { + t.Fatalf("expected %d values, got %d", len(keys), len(values)) + } + for _, k := range keys { + if v, ok := values[k]; !ok || v != k+"_v" { + t.Fatalf("unexpected value for %s: %q, ok=%v", k, v, ok) + } + } + + // second batch should hit cache + values, err = l.GetBatch(keys, func(keys []string) (map[string]string, error) { + t.Fatalf("should not fetch on cache hit") + return nil, nil + }) + if err != nil { + t.Fatalf("unexpected error on cache hit: %v", err) + } + for _, k := range keys { + if v, ok := values[k]; !ok || v != k+"_v" { + t.Fatalf("unexpected cached value for %s: %q, ok=%v", k, v, ok) + } + } +} diff --git a/pkg/localcache/init.go b/pkg/localcache/init.go index ad339da7c..d0bccaa7e 100644 --- a/pkg/localcache/init.go +++ b/pkg/localcache/init.go @@ -33,10 +33,6 @@ func InitLocalCache(localCache *config.LocalCache) { Local config.CacheConfig Keys []string }{ - { - Local: localCache.Auth, - Keys: []string{cachekey.UidPidToken}, - }, { Local: localCache.User, Keys: []string{cachekey.UserInfoKey, cachekey.UserGlobalRecvMsgOptKey}, diff --git a/pkg/localcache/lru/lru_expiration.go b/pkg/localcache/lru/lru_expiration.go index df6bacbf4..4197cacec 100644 --- a/pkg/localcache/lru/lru_expiration.go +++ b/pkg/localcache/lru/lru_expiration.go @@ -52,8 +52,53 @@ type ExpirationLRU[K comparable, V any] struct { } func (x *ExpirationLRU[K, V]) GetBatch(keys []K, fetch func(keys []K) (map[K]V, error)) (map[K]V, error) { - //TODO implement me - panic("implement me") + var ( + err error + results = make(map[K]V) + misses = make([]K, 0, len(keys)) + ) + + for _, key := range keys { + x.lock.Lock() + v, ok := x.core.Get(key) + x.lock.Unlock() + if ok { + x.target.IncrGetHit() + v.lock.RLock() + results[key] = v.value + if v.err != nil && err == nil { + err = v.err + } + v.lock.RUnlock() + continue + } + misses = append(misses, key) + } + + if len(misses) == 0 { + return results, err + } + + fetchValues, fetchErr := fetch(misses) + if fetchErr != nil && err == nil { + err = fetchErr + } + + for key, val := range fetchValues { + results[key] = val + if fetchErr != nil { + x.target.IncrGetFailed() + continue + } + x.target.IncrGetSuccess() + item := &expirationLruItem[V]{value: val} + x.lock.Lock() + x.core.Add(key, item) + x.lock.Unlock() + } + + // any keys not returned from fetch remain absent (no cache write) + return results, err } func (x *ExpirationLRU[K, V]) Get(key K, fetch func() (V, error)) (V, error) { diff --git a/pkg/localcache/lru/lru_slot.go b/pkg/localcache/lru/lru_slot.go index 14ee3b50f..077219b75 100644 --- a/pkg/localcache/lru/lru_slot.go +++ b/pkg/localcache/lru/lru_slot.go @@ -35,7 +35,7 @@ type slotLRU[K comparable, V any] struct { func (x *slotLRU[K, V]) GetBatch(keys []K, fetch func(keys []K) (map[K]V, error)) (map[K]V, error) { var ( slotKeys = make(map[uint64][]K) - kVs = make(map[K]V) + vs = make(map[K]V) ) for _, k := range keys { @@ -49,10 +49,10 @@ func (x *slotLRU[K, V]) GetBatch(keys []K, fetch func(keys []K) (map[K]V, error) return nil, err } for key, value := range batches { - kVs[key] = value + vs[key] = value } } - return kVs, nil + return vs, nil } func (x *slotLRU[K, V]) getIndex(k K) uint64 { From a1dd79a4592f1be55ad0c96fece250df7e185002 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 19 Dec 2025 15:11:41 +0800 Subject: [PATCH 05/19] refactor: replace LongConn with ClientConn interface and simplify message handling --- internal/msggateway/client.go | 163 ++------------------ internal/msggateway/client_conn.go | 229 +++++++++++++++++++++++++++++ internal/msggateway/context.go | 216 ++++++++++++++++----------- internal/msggateway/long_conn.go | 179 ---------------------- internal/msggateway/ws_server.go | 108 +++++++------- 5 files changed, 427 insertions(+), 468 deletions(-) create mode 100644 internal/msggateway/client_conn.go delete mode 100644 internal/msggateway/long_conn.go diff --git a/internal/msggateway/client.go b/internal/msggateway/client.go index 7b9f4bc0e..46da524d3 100644 --- a/internal/msggateway/client.go +++ b/internal/msggateway/client.go @@ -16,7 +16,6 @@ package msggateway import ( "context" - "encoding/json" "fmt" "sync" "sync/atomic" @@ -31,7 +30,6 @@ import ( "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/mcontext" - "github.com/openimsdk/tools/utils/stringutil" ) var ( @@ -64,13 +62,12 @@ type PingPongHandler func(string) error type Client struct { w *sync.Mutex - conn LongConn + conn ClientConn PlatformID int `json:"platformID"` IsCompress bool `json:"isCompress"` UserID string `json:"userID"` IsBackground bool `json:"isBackground"` SDKType string `json:"sdkType"` - SDKVersion string `json:"sdkVersion"` Encoder Encoder ctx *UserConnContext longConnServer LongConnServer @@ -84,10 +81,10 @@ type Client struct { } // ResetClient updates the client's state with new connection and context information. -func (c *Client) ResetClient(ctx *UserConnContext, conn LongConn, longConnServer LongConnServer) { +func (c *Client) ResetClient(ctx *UserConnContext, conn ClientConn, longConnServer LongConnServer) { c.w = new(sync.Mutex) c.conn = conn - c.PlatformID = stringutil.StringToInt(ctx.GetPlatformID()) + c.PlatformID = ctx.GetPlatformID() c.IsCompress = ctx.GetCompression() c.IsBackground = ctx.GetBackground() c.UserID = ctx.GetUserID() @@ -98,7 +95,6 @@ func (c *Client) ResetClient(ctx *UserConnContext, conn LongConn, longConnServer c.closedErr = nil c.token = ctx.GetToken() c.SDKType = ctx.GetSDKType() - c.SDKVersion = ctx.GetSDKVersion() c.hbCtx, c.hbCancel = context.WithCancel(c.ctx) c.subLock = new(sync.Mutex) if c.subUserIDs != nil { @@ -112,22 +108,6 @@ func (c *Client) ResetClient(ctx *UserConnContext, conn LongConn, longConnServer c.subUserIDs = make(map[string]struct{}) } -func (c *Client) pingHandler(appData string) error { - if err := c.conn.SetReadDeadline(pongWait); err != nil { - return err - } - - log.ZDebug(c.ctx, "ping Handler Success.", "appData", appData) - return c.writePongMsg(appData) -} - -func (c *Client) pongHandler(_ string) error { - if err := c.conn.SetReadDeadline(pongWait); err != nil { - return err - } - return nil -} - // readMessage continuously reads messages from the connection. func (c *Client) readMessage() { defer func() { @@ -138,52 +118,25 @@ func (c *Client) readMessage() { c.close() }() - c.conn.SetReadLimit(maxMessageSize) - _ = c.conn.SetReadDeadline(pongWait) - c.conn.SetPongHandler(c.pongHandler) - c.conn.SetPingHandler(c.pingHandler) - c.activeHeartbeat(c.hbCtx) - for { log.ZDebug(c.ctx, "readMessage") - messageType, message, returnErr := c.conn.ReadMessage() + message, returnErr := c.conn.ReadMessage() if returnErr != nil { - log.ZWarn(c.ctx, "readMessage", returnErr, "messageType", messageType) + log.ZWarn(c.ctx, "readMessage", returnErr) c.closedErr = returnErr return } - log.ZDebug(c.ctx, "readMessage", "messageType", messageType) if c.closed.Load() { // The scenario where the connection has just been closed, but the coroutine has not exited c.closedErr = ErrConnClosed return } - switch messageType { - case MessageBinary: - _ = c.conn.SetReadDeadline(pongWait) - parseDataErr := c.handleMessage(message) - if parseDataErr != nil { - c.closedErr = parseDataErr - return - } - case MessageText: - _ = c.conn.SetReadDeadline(pongWait) - parseDataErr := c.handlerTextMessage(message) - if parseDataErr != nil { - c.closedErr = parseDataErr - return - } - case PingMessage: - err := c.writePongMsg("") - log.ZError(c.ctx, "writePongMsg", err) - - case CloseMessage: - c.closedErr = ErrClientClosed + parseDataErr := c.handleMessage(message) + if parseDataErr != nil { + c.closedErr = parseDataErr return - - default: } } } @@ -358,109 +311,13 @@ func (c *Client) writeBinaryMsg(resp Resp) error { c.w.Lock() defer c.w.Unlock() - err = c.conn.SetWriteDeadline(writeWait) - if err != nil { - return err - } - if c.IsCompress { resultBuf, compressErr := c.longConnServer.CompressWithPool(encodedBuf) if compressErr != nil { return compressErr } - return c.conn.WriteMessage(MessageBinary, resultBuf) + return c.conn.WriteMessage(resultBuf) } - return c.conn.WriteMessage(MessageBinary, encodedBuf) -} - -// Actively initiate Heartbeat when platform in Web. -func (c *Client) activeHeartbeat(ctx context.Context) { - if c.PlatformID == constant.WebPlatformID { - go func() { - defer func() { - if r := recover(); r != nil { - log.ZPanic(ctx, "activeHeartbeat Panic", errs.ErrPanic(r)) - } - }() - log.ZDebug(ctx, "server initiative send heartbeat start.") - ticker := time.NewTicker(pingPeriod) - defer ticker.Stop() - - for { - select { - case <-ticker.C: - if err := c.writePingMsg(); err != nil { - log.ZWarn(c.ctx, "send Ping Message error.", err) - return - } - case <-c.hbCtx.Done(): - return - } - } - }() - } -} -func (c *Client) writePingMsg() error { - if c.closed.Load() { - return nil - } - - c.w.Lock() - defer c.w.Unlock() - - err := c.conn.SetWriteDeadline(writeWait) - if err != nil { - return err - } - - return c.conn.WriteMessage(PingMessage, nil) -} - -func (c *Client) writePongMsg(appData string) error { - log.ZDebug(c.ctx, "write Pong Msg in Server", "appData", appData) - if c.closed.Load() { - log.ZWarn(c.ctx, "is closed in server", nil, "appdata", appData, "closed err", c.closedErr) - return nil - } - - c.w.Lock() - defer c.w.Unlock() - - err := c.conn.SetWriteDeadline(writeWait) - if err != nil { - log.ZWarn(c.ctx, "SetWriteDeadline in Server have error", errs.Wrap(err), "writeWait", writeWait, "appData", appData) - return errs.Wrap(err) - } - err = c.conn.WriteMessage(PongMessage, []byte(appData)) - if err != nil { - log.ZWarn(c.ctx, "Write Message have error", errs.Wrap(err), "Pong msg", PongMessage) - } - - return errs.Wrap(err) -} - -func (c *Client) handlerTextMessage(b []byte) error { - var msg TextMessage - if err := json.Unmarshal(b, &msg); err != nil { - return err - } - switch msg.Type { - case TextPong: - return nil - case TextPing: - msg.Type = TextPong - msgData, err := json.Marshal(msg) - if err != nil { - return err - } - c.w.Lock() - defer c.w.Unlock() - if err := c.conn.SetWriteDeadline(writeWait); err != nil { - return err - } - return c.conn.WriteMessage(MessageText, msgData) - default: - return fmt.Errorf("not support message type %s", msg.Type) - } + return c.conn.WriteMessage(encodedBuf) } diff --git a/internal/msggateway/client_conn.go b/internal/msggateway/client_conn.go new file mode 100644 index 000000000..15a0d8c07 --- /dev/null +++ b/internal/msggateway/client_conn.go @@ -0,0 +1,229 @@ +package msggateway + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sync/atomic" + "time" + + "github.com/gorilla/websocket" + + "github.com/openimsdk/tools/log" +) + +var ErrWriteFull = fmt.Errorf("websocket write buffer full,close connection") + +type ClientConn interface { + ReadMessage() ([]byte, error) + WriteMessage(message []byte) error + Close() error +} + +type websocketMessage struct { + MessageType int + Data []byte +} + +func NewWebSocketClientConn(conn *websocket.Conn, readLimit int64, readTimeout time.Duration, pingInterval time.Duration) ClientConn { + c := &websocketClientConn{ + readTimeout: readTimeout, + conn: conn, + writer: make(chan *websocketMessage, 256), + done: make(chan struct{}), + } + if readLimit > 0 { + c.conn.SetReadLimit(readLimit) + } + c.conn.SetPingHandler(c.pingHandler) + c.conn.SetPongHandler(c.pongHandler) + + go c.loopSend() + if pingInterval > 0 { + go c.doPing(pingInterval) + } + return c +} + +type websocketClientConn struct { + readTimeout time.Duration + conn *websocket.Conn + writer chan *websocketMessage + done chan struct{} + err atomic.Pointer[error] +} + +func (c *websocketClientConn) ReadMessage() ([]byte, error) { + buf, err := c.readMessage() + if err != nil { + return nil, c.closeBy(fmt.Errorf("read message %w", err)) + } + return buf, nil +} + +func (c *websocketClientConn) WriteMessage(message []byte) error { + return c.writeMessage(websocket.BinaryMessage, message) +} + +func (c *websocketClientConn) Close() error { + return c.closeBy(fmt.Errorf("websocket connection closed")) +} + +func (c *websocketClientConn) closeBy(err error) error { + if !c.err.CompareAndSwap(nil, &err) { + return *c.err.Load() + } + close(c.done) + log.ZWarn(context.Background(), "websocket connection closed", err, "remoteAddr", c.conn.RemoteAddr(), + "chan length", len(c.writer)) + return err +} + +func (c *websocketClientConn) writeMessage(messageType int, data []byte) error { + if errPtr := c.err.Load(); errPtr != nil { + return *errPtr + } + select { + case c.writer <- &websocketMessage{MessageType: messageType, Data: data}: + return nil + default: + return c.closeBy(ErrWriteFull) + } +} + +func (c *websocketClientConn) loopSend() { + defer func() { + _ = c.conn.Close() + }() + var err error + for { + select { + case <-c.done: + for { + select { + case msg := <-c.writer: + switch msg.MessageType { + case websocket.TextMessage, websocket.BinaryMessage: + err = c.conn.WriteMessage(msg.MessageType, msg.Data) + default: + err = c.conn.WriteControl(msg.MessageType, msg.Data, time.Time{}) + } + if err != nil { + _ = c.closeBy(err) + return + } + default: + return + } + } + case msg := <-c.writer: + switch msg.MessageType { + case websocket.TextMessage, websocket.BinaryMessage: + err = c.conn.WriteMessage(msg.MessageType, msg.Data) + default: + err = c.conn.WriteControl(msg.MessageType, msg.Data, time.Time{}) + } + if err != nil { + _ = c.closeBy(err) + return + } + } + } +} + +func (c *websocketClientConn) setReadDeadline() error { + deadline := time.Now().Add(c.readTimeout) + return c.conn.SetReadDeadline(deadline) +} + +func (c *websocketClientConn) readMessage() ([]byte, error) { + for { + if err := c.setReadDeadline(); err != nil { + return nil, err + } + messageType, buf, err := c.conn.ReadMessage() + if err != nil { + return nil, err + } + switch messageType { + case websocket.BinaryMessage: + return buf, nil + case websocket.TextMessage: + if err := c.onReadTextMessage(buf); err != nil { + return nil, err + } + case websocket.PingMessage: + if err := c.pingHandler(string(buf)); err != nil { + return nil, err + } + case websocket.PongMessage: + if err := c.pongHandler(string(buf)); err != nil { + return nil, err + } + case websocket.CloseMessage: + if len(buf) == 0 { + return nil, errors.New("websocket connection closed by peer") + } + return nil, fmt.Errorf("websocket connection closed by peer, data %s", string(buf)) + default: + return nil, fmt.Errorf("unknown websocket message type %d", messageType) + } + } +} + +func (c *websocketClientConn) onReadTextMessage(buf []byte) error { + var msg struct { + Type string `json:"type"` + Body json.RawMessage `json:"body"` + } + if err := json.Unmarshal(buf, &msg); err != nil { + return err + } + switch msg.Type { + case TextPong: + return nil + case TextPing: + msg.Type = TextPong + msgData, err := json.Marshal(msg) + if err != nil { + return err + } + return c.writeMessage(websocket.TextMessage, msgData) + default: + return fmt.Errorf("not support text message type %s", msg.Type) + } +} + +func (c *websocketClientConn) pingHandler(appData string) error { + log.ZDebug(context.Background(), "ping handler recv ping", "remoteAddr", c.conn.RemoteAddr(), "appData", appData) + if err := c.setReadDeadline(); err != nil { + return err + } + err := c.conn.WriteControl(websocket.PongMessage, []byte(appData), time.Now().Add(time.Second*1)) + if err != nil { + log.ZWarn(context.Background(), "ping handler write pong error", err, "remoteAddr", c.conn.RemoteAddr(), "appData", appData) + } + log.ZDebug(context.Background(), "ping handler write pong success", "remoteAddr", c.conn.RemoteAddr(), "appData", appData) + return nil +} + +func (c *websocketClientConn) pongHandler(string) error { + return nil +} + +func (c *websocketClientConn) doPing(d time.Duration) { + ticker := time.NewTicker(d) + defer ticker.Stop() + for { + select { + case <-c.done: + return + case <-ticker.C: + if err := c.writeMessage(websocket.PingMessage, nil); err != nil { + _ = c.closeBy(fmt.Errorf("send ping %w", err)) + return + } + } + } +} diff --git a/internal/msggateway/context.go b/internal/msggateway/context.go index 37b5a7cdc..6883c22a3 100644 --- a/internal/msggateway/context.go +++ b/internal/msggateway/context.go @@ -15,6 +15,8 @@ package msggateway import ( + "encoding/base64" + "encoding/json" "net/http" "net/url" "strconv" @@ -24,10 +26,20 @@ import ( "github.com/openimsdk/protocol/constant" "github.com/openimsdk/tools/utils/encrypt" - "github.com/openimsdk/tools/utils/stringutil" "github.com/openimsdk/tools/utils/timeutil" ) +type UserConnContextInfo struct { + Token string `json:"token"` + UserID string `json:"userID"` + PlatformID int `json:"platformID"` + OperationID string `json:"operationID"` + Compression string `json:"compression"` + SDKType string `json:"sdkType"` + SendResponse bool `json:"sendResponse"` + Background bool `json:"background"` +} + type UserConnContext struct { RespWriter http.ResponseWriter Req *http.Request @@ -35,6 +47,7 @@ type UserConnContext struct { Method string RemoteAddr string ConnID string + info *UserConnContextInfo } func (c *UserConnContext) Deadline() (deadline time.Time, ok bool) { @@ -58,7 +71,7 @@ func (c *UserConnContext) Value(key any) any { case constant.ConnID: return c.GetConnID() case constant.OpUserPlatform: - return constant.PlatformIDToName(stringutil.StringToInt(c.GetPlatformID())) + return c.GetPlatformID() case constant.RemoteAddr: return c.RemoteAddr default: @@ -83,30 +96,91 @@ func newContext(respWriter http.ResponseWriter, req *http.Request) *UserConnCont func newTempContext() *UserConnContext { return &UserConnContext{ - Req: &http.Request{URL: &url.URL{}}, + Req: &http.Request{URL: &url.URL{}}, + info: &UserConnContextInfo{}, } } +func (c *UserConnContext) ParseEssentialArgs() error { + query := c.Req.URL.Query() + if data := query.Get("v"); data != "" { + return c.parseByJson(data) + } else { + return c.parseByQuery(query, c.Req.Header) + } +} + +func (c *UserConnContext) parseByQuery(query url.Values, header http.Header) error { + info := UserConnContextInfo{ + Token: query.Get(Token), + UserID: query.Get(WsUserID), + OperationID: query.Get(OperationID), + Compression: query.Get(Compression), + SDKType: query.Get(SDKType), + } + platformID, err := strconv.Atoi(query.Get(PlatformID)) + if err != nil { + return servererrs.ErrConnArgsErr.WrapMsg("platformID is not int") + } + info.PlatformID = platformID + if val := query.Get(SendResponse); val != "" { + ok, err := strconv.ParseBool(val) + if err != nil { + return servererrs.ErrConnArgsErr.WrapMsg("isMsgResp is not bool") + } + info.SendResponse = ok + } + if info.Compression == "" { + info.Compression = header.Get(Compression) + } + background, err := strconv.ParseBool(query.Get(BackgroundStatus)) + if err != nil { + return err + } + info.Background = background + return c.checkInfo(&info) +} + +func (c *UserConnContext) parseByJson(data string) error { + reqInfo, err := base64.RawURLEncoding.DecodeString(data) + if err != nil { + return servererrs.ErrConnArgsErr.WrapMsg("data is not base64") + } + var info UserConnContextInfo + if err := json.Unmarshal(reqInfo, &info); err != nil { + return servererrs.ErrConnArgsErr.WrapMsg("data is not json", "info", err.Error()) + } + return c.checkInfo(&info) +} + +func (c *UserConnContext) checkInfo(info *UserConnContextInfo) error { + if info.OperationID == "" { + return servererrs.ErrConnArgsErr.WrapMsg("operationID is empty") + } + if info.Token == "" { + return servererrs.ErrConnArgsErr.WrapMsg("token is empty") + } + if info.UserID == "" { + return servererrs.ErrConnArgsErr.WrapMsg("sendID is empty") + } + if _, ok := constant.PlatformID2Name[info.PlatformID]; !ok { + return servererrs.ErrConnArgsErr.WrapMsg("platformID is invalid") + } + switch info.SDKType { + case "": + info.SDKType = GoSDK + case GoSDK, JsSDK: + default: + return servererrs.ErrConnArgsErr.WrapMsg("sdkType is invalid") + } + c.info = info + return nil +} + func (c *UserConnContext) GetRemoteAddr() string { return c.RemoteAddr } -func (c *UserConnContext) Query(key string) (string, bool) { - var value string - if value = c.Req.URL.Query().Get(key); value == "" { - return value, false - } - return value, true -} - -func (c *UserConnContext) GetHeader(key string) (string, bool) { - var value string - if value = c.Req.Header.Get(key); value == "" { - return value, false - } - return value, true -} - func (c *UserConnContext) SetHeader(key, value string) { c.RespWriter.Header().Set(key, value) } @@ -120,97 +194,69 @@ func (c *UserConnContext) GetConnID() string { } func (c *UserConnContext) GetUserID() string { - return c.Req.URL.Query().Get(WsUserID) + if c == nil || c.info == nil { + return "" + } + return c.info.UserID } -func (c *UserConnContext) GetPlatformID() string { - return c.Req.URL.Query().Get(PlatformID) +func (c *UserConnContext) GetPlatformID() int { + if c == nil || c.info == nil { + return 0 + } + return c.info.PlatformID } func (c *UserConnContext) GetOperationID() string { - return c.Req.URL.Query().Get(OperationID) + if c == nil || c.info == nil { + return "" + } + return c.info.OperationID } func (c *UserConnContext) SetOperationID(operationID string) { - values := c.Req.URL.Query() - values.Set(OperationID, operationID) - c.Req.URL.RawQuery = values.Encode() + if c.info == nil { + c.info = &UserConnContextInfo{} + } + c.info.OperationID = operationID } func (c *UserConnContext) GetToken() string { - return c.Req.URL.Query().Get(Token) -} - -func (c *UserConnContext) GetSDKVersion() string { - return c.Req.URL.Query().Get(SDKVersion) + if c == nil || c.info == nil { + return "" + } + return c.info.Token } func (c *UserConnContext) GetCompression() bool { - compression, exists := c.Query(Compression) - if exists && compression == GzipCompressionProtocol { - return true - } else { - compression, exists := c.GetHeader(Compression) - if exists && compression == GzipCompressionProtocol { - return true - } - } - return false + return c != nil && c.info != nil && c.info.Compression == GzipCompressionProtocol } func (c *UserConnContext) GetSDKType() string { - sdkType := c.Req.URL.Query().Get(SDKType) - if sdkType == "" { - sdkType = GoSDK + if c == nil || c.info == nil { + return GoSDK + } + switch c.info.SDKType { + case "", GoSDK: + return GoSDK + case JsSDK: + return JsSDK + default: + return "" } - return sdkType } func (c *UserConnContext) ShouldSendResp() bool { - errResp, exists := c.Query(SendResponse) - if exists { - b, err := strconv.ParseBool(errResp) - if err != nil { - return false - } else { - return b - } - } - return false + return c != nil && c.info != nil && c.info.SendResponse } func (c *UserConnContext) SetToken(token string) { - c.Req.URL.RawQuery = Token + "=" + token + if c.info == nil { + c.info = &UserConnContextInfo{} + } + c.info.Token = token } func (c *UserConnContext) GetBackground() bool { - b, err := strconv.ParseBool(c.Req.URL.Query().Get(BackgroundStatus)) - if err != nil { - return false - } - return b -} -func (c *UserConnContext) ParseEssentialArgs() error { - _, exists := c.Query(Token) - if !exists { - return servererrs.ErrConnArgsErr.WrapMsg("token is empty") - } - _, exists = c.Query(WsUserID) - if !exists { - return servererrs.ErrConnArgsErr.WrapMsg("sendID is empty") - } - platformIDStr, exists := c.Query(PlatformID) - if !exists { - return servererrs.ErrConnArgsErr.WrapMsg("platformID is empty") - } - _, err := strconv.Atoi(platformIDStr) - if err != nil { - return servererrs.ErrConnArgsErr.WrapMsg("platformID is not int") - } - switch sdkType, _ := c.Query(SDKType); sdkType { - case "", GoSDK, JsSDK: - default: - return servererrs.ErrConnArgsErr.WrapMsg("sdkType is not go or js") - } - return nil + return c != nil && c.info != nil && c.info.Background } diff --git a/internal/msggateway/long_conn.go b/internal/msggateway/long_conn.go deleted file mode 100644 index c1b3e27c9..000000000 --- a/internal/msggateway/long_conn.go +++ /dev/null @@ -1,179 +0,0 @@ -// Copyright © 2023 OpenIM. All rights reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package msggateway - -import ( - "encoding/json" - "net/http" - "time" - - "github.com/openimsdk/tools/apiresp" - - "github.com/gorilla/websocket" - "github.com/openimsdk/tools/errs" -) - -type LongConn interface { - // Close this connection - Close() error - // WriteMessage Write message to connection,messageType means data type,can be set binary(2) and text(1). - WriteMessage(messageType int, message []byte) error - // ReadMessage Read message from connection. - ReadMessage() (int, []byte, error) - // SetReadDeadline sets the read deadline on the underlying network connection, - // after a read has timed out, will return an error. - SetReadDeadline(timeout time.Duration) error - // SetWriteDeadline sets to write deadline when send message,when read has timed out,will return error. - SetWriteDeadline(timeout time.Duration) error - // Dial Try to dial a connection,url must set auth args,header can control compress data - Dial(urlStr string, requestHeader http.Header) (*http.Response, error) - // IsNil Whether the connection of the current long connection is nil - IsNil() bool - // SetConnNil Set the connection of the current long connection to nil - SetConnNil() - // SetReadLimit sets the maximum size for a message read from the peer.bytes - SetReadLimit(limit int64) - SetPongHandler(handler PingPongHandler) - SetPingHandler(handler PingPongHandler) - // GenerateLongConn Check the connection of the current and when it was sent are the same - GenerateLongConn(w http.ResponseWriter, r *http.Request) error -} -type GWebSocket struct { - protocolType int - conn *websocket.Conn - handshakeTimeout time.Duration - writeBufferSize int -} - -func newGWebSocket(protocolType int, handshakeTimeout time.Duration, wbs int) *GWebSocket { - return &GWebSocket{protocolType: protocolType, handshakeTimeout: handshakeTimeout, writeBufferSize: wbs} -} - -func (d *GWebSocket) Close() error { - return d.conn.Close() -} - -func (d *GWebSocket) GenerateLongConn(w http.ResponseWriter, r *http.Request) error { - upgrader := &websocket.Upgrader{ - HandshakeTimeout: d.handshakeTimeout, - CheckOrigin: func(r *http.Request) bool { return true }, - } - if d.writeBufferSize > 0 { // default is 4kb. - upgrader.WriteBufferSize = d.writeBufferSize - } - - conn, err := upgrader.Upgrade(w, r, nil) - if err != nil { - // The upgrader.Upgrade method usually returns enough error messages to diagnose problems that may occur during the upgrade - return errs.WrapMsg(err, "GenerateLongConn: WebSocket upgrade failed") - } - d.conn = conn - return nil -} - -func (d *GWebSocket) WriteMessage(messageType int, message []byte) error { - // d.setSendConn(d.conn) - return d.conn.WriteMessage(messageType, message) -} - -// func (d *GWebSocket) setSendConn(sendConn *websocket.Conn) { -// d.sendConn = sendConn -//} - -func (d *GWebSocket) ReadMessage() (int, []byte, error) { - return d.conn.ReadMessage() -} - -func (d *GWebSocket) SetReadDeadline(timeout time.Duration) error { - return d.conn.SetReadDeadline(time.Now().Add(timeout)) -} - -func (d *GWebSocket) SetWriteDeadline(timeout time.Duration) error { - if timeout <= 0 { - return errs.New("timeout must be greater than 0") - } - - // TODO SetWriteDeadline Future add error handling - if err := d.conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil { - return errs.WrapMsg(err, "GWebSocket.SetWriteDeadline failed") - } - return nil -} - -func (d *GWebSocket) Dial(urlStr string, requestHeader http.Header) (*http.Response, error) { - conn, httpResp, err := websocket.DefaultDialer.Dial(urlStr, requestHeader) - if err != nil { - return httpResp, errs.WrapMsg(err, "GWebSocket.Dial failed", "url", urlStr) - } - d.conn = conn - return httpResp, nil -} - -func (d *GWebSocket) IsNil() bool { - return d.conn == nil - // - // if d.conn != nil { - // return false - // } - // return true -} - -func (d *GWebSocket) SetConnNil() { - d.conn = nil -} - -func (d *GWebSocket) SetReadLimit(limit int64) { - d.conn.SetReadLimit(limit) -} - -func (d *GWebSocket) SetPongHandler(handler PingPongHandler) { - d.conn.SetPongHandler(handler) -} - -func (d *GWebSocket) SetPingHandler(handler PingPongHandler) { - d.conn.SetPingHandler(handler) -} - -func (d *GWebSocket) RespondWithError(err error, w http.ResponseWriter, r *http.Request) error { - if err := d.GenerateLongConn(w, r); err != nil { - return err - } - data, err := json.Marshal(apiresp.ParseError(err)) - if err != nil { - _ = d.Close() - return errs.WrapMsg(err, "json marshal failed") - } - - if err := d.WriteMessage(MessageText, data); err != nil { - _ = d.Close() - return errs.WrapMsg(err, "WriteMessage failed") - } - _ = d.Close() - return nil -} - -func (d *GWebSocket) RespondWithSuccess() error { - data, err := json.Marshal(apiresp.ParseError(nil)) - if err != nil { - _ = d.Close() - return errs.WrapMsg(err, "json marshal failed") - } - - if err := d.WriteMessage(MessageText, data); err != nil { - _ = d.Close() - return errs.WrapMsg(err, "WriteMessage failed") - } - return nil -} diff --git a/internal/msggateway/ws_server.go b/internal/msggateway/ws_server.go index d490cc8b9..0f7e1f8e6 100644 --- a/internal/msggateway/ws_server.go +++ b/internal/msggateway/ws_server.go @@ -2,18 +2,20 @@ package msggateway import ( "context" + "encoding/json" "fmt" "net/http" "sync" "sync/atomic" "time" + "github.com/gorilla/websocket" "github.com/openimsdk/open-im-server/v3/pkg/rpcli" + "github.com/openimsdk/tools/apiresp" "github.com/openimsdk/open-im-server/v3/pkg/common/webhook" "github.com/openimsdk/open-im-server/v3/pkg/rpccache" pbAuth "github.com/openimsdk/protocol/auth" - "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/mcontext" "github.com/go-playground/validator/v10" @@ -23,10 +25,11 @@ import ( "github.com/openimsdk/protocol/msggateway" "github.com/openimsdk/tools/discovery" "github.com/openimsdk/tools/log" - "github.com/openimsdk/tools/utils/stringutil" "golang.org/x/sync/errgroup" ) +var wsSuccessResponse, _ = json.Marshal(&apiresp.ApiResponse{}) + type LongConnServer interface { Run(ctx context.Context) error wsHandler(w http.ResponseWriter, r *http.Request) @@ -43,6 +46,7 @@ type LongConnServer interface { } type WsServer struct { + websocket *websocket.Upgrader msgGatewayConfig *Config port int wsMaxConnNum int64 @@ -136,9 +140,13 @@ func NewWsServer(msgGatewayConfig *Config, opts ...Option) *WsServer { o(&config) } //userRpcClient := rpcclient.NewUserRpcClient(client, config.Discovery.RpcService.User, config.Share.IMAdminUser) - + upgrader := &websocket.Upgrader{ + HandshakeTimeout: config.handshakeTimeout, + CheckOrigin: func(r *http.Request) bool { return true }, + } v := validator.New() return &WsServer{ + websocket: upgrader, msgGatewayConfig: msgGatewayConfig, port: config.port, wsMaxConnNum: config.maxConnNum, @@ -260,8 +268,7 @@ func (ws *WsServer) registerClient(client *Client) { ) oldClients, userOK, clientOK = ws.clients.Get(client.UserID, client.PlatformID) - log.ZInfo(client.ctx, "registerClient", "userID", client.UserID, "platformID", client.PlatformID, - "sdkVersion", client.SDKVersion) + log.ZInfo(client.ctx, "registerClient", "userID", client.UserID, "platformID", client.PlatformID) if !userOK { ws.clients.Set(client.UserID, client) @@ -448,7 +455,7 @@ func (ws *WsServer) unregisterClient(client *Client) { // validateRespWithRequest checks if the response matches the expected userID and platformID. func (ws *WsServer) validateRespWithRequest(ctx *UserConnContext, resp *pbAuth.ParseTokenResp) error { userID := ctx.GetUserID() - platformID := stringutil.StringToInt32(ctx.GetPlatformID()) + platformID := int32(ctx.GetPlatformID()) if resp.UserID != userID { return servererrs.ErrTokenInvalid.WrapMsg(fmt.Sprintf("token uid %s != userID %s", resp.UserID, userID)) } @@ -458,19 +465,37 @@ func (ws *WsServer) validateRespWithRequest(ctx *UserConnContext, resp *pbAuth.P return nil } +func (ws *WsServer) handlerError(ctx *UserConnContext, w http.ResponseWriter, r *http.Request, err error) { + if !ctx.ShouldSendResp() { + httpError(ctx, err) + return + } + // the browser cannot get the response of upgrade failure + data, err := json.Marshal(apiresp.ParseError(err)) + if err != nil { + log.ZError(ctx, "json marshal failed", err) + return + } + conn, upgradeErr := ws.websocket.Upgrade(w, r, nil) + if upgradeErr != nil { + log.ZWarn(ctx, "websocket upgrade failed", upgradeErr, "respErr", err, "resp", string(data)) + return + } + defer conn.Close() + if err := conn.WriteMessage(websocket.TextMessage, data); err != nil { + log.ZWarn(ctx, "WriteMessage failed", err, "respErr", err, "resp", string(data)) + return + } +} + func (ws *WsServer) wsHandler(w http.ResponseWriter, r *http.Request) { // Create a new connection context connContext := newContext(w, r) - if !ws.ready.Load() { - httpError(connContext, errs.New("ws server not ready")) - return - } - // Check if the current number of online user connections exceeds the maximum limit if ws.onlineUserConnNum.Load() >= ws.wsMaxConnNum { // If it exceeds the maximum connection number, return an error via HTTP and stop processing - httpError(connContext, servererrs.ErrConnOverMaxNumLimit.WrapMsg("over max conn num limit")) + ws.handlerError(connContext, w, r, servererrs.ErrConnOverMaxNumLimit.WrapMsg("over max conn num limit")) return } @@ -478,31 +503,14 @@ func (ws *WsServer) wsHandler(w http.ResponseWriter, r *http.Request) { err := connContext.ParseEssentialArgs() if err != nil { // If there's an error during parsing, return an error via HTTP and stop processing - - httpError(connContext, err) - return - } - - if ws.authClient == nil { - httpError(connContext, errs.New("auth client is not initialized")) + ws.handlerError(connContext, w, r, err) return } // Call the authentication client to parse the Token obtained from the context resp, err := ws.authClient.ParseToken(connContext, connContext.GetToken()) if err != nil { - // If there's an error parsing the Token, decide whether to send the error message via WebSocket based on the context flag - shouldSendError := connContext.ShouldSendResp() - if shouldSendError { - // Create a WebSocket connection object and attempt to send the error message via WebSocket - wsLongConn := newGWebSocket(WebSocket, ws.handshakeTimeout, ws.writeBufferSize) - if err := wsLongConn.RespondWithError(err, w, r); err == nil { - // If the error message is successfully sent via WebSocket, stop processing - return - } - } - // If sending via WebSocket is not required or fails, return the error via HTTP and stop processing - httpError(connContext, err) + ws.handlerError(connContext, w, r, err) return } @@ -510,32 +518,30 @@ func (ws *WsServer) wsHandler(w http.ResponseWriter, r *http.Request) { err = ws.validateRespWithRequest(connContext, resp) if err != nil { // If validation fails, return an error via HTTP and stop processing - httpError(connContext, err) + ws.handlerError(connContext, w, r, err) return } - - log.ZDebug(connContext, "new conn", "token", connContext.GetToken()) - // Create a WebSocket long connection object - wsLongConn := newGWebSocket(WebSocket, ws.handshakeTimeout, ws.writeBufferSize) - if err := wsLongConn.GenerateLongConn(w, r); err != nil { - //If the creation of the long connection fails, the error is handled internally during the handshake process. - log.ZWarn(connContext, "long connection fails", err) + conn, err := ws.websocket.Upgrade(w, r, nil) + if err != nil { + log.ZWarn(connContext, "websocket upgrade failed", err) return - } else { - // Check if a normal response should be sent via WebSocket - shouldSendSuccessResp := connContext.ShouldSendResp() - if shouldSendSuccessResp { - // Attempt to send a success message through WebSocket - if err := wsLongConn.RespondWithSuccess(); err != nil { - // If the success message is successfully sent, end further processing - return - } + } + if connContext.ShouldSendResp() { + if err := conn.WriteMessage(websocket.TextMessage, wsSuccessResponse); err != nil { + log.ZWarn(connContext, "WriteMessage first response", err) + return } } - // Retrieve a client object from the client pool, reset its state, and associate it with the current WebSocket long connection - client := ws.clientPool.Get().(*Client) - client.ResetClient(connContext, wsLongConn, ws) + log.ZDebug(connContext, "new conn", "token", connContext.GetToken()) + + var pingInterval time.Duration + if connContext.GetPlatformID() == constant.WebPlatformID { + pingInterval = pingPeriod + } + + client := new(Client) + client.ResetClient(connContext, NewWebSocketClientConn(conn, maxMessageSize, pongWait, pingInterval), ws) // Register the client with the server and start message processing ws.registerChan <- client From 9da7db2ac25e93a937a3f4dc6d80771c248f60be Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 19 Dec 2025 15:53:17 +0800 Subject: [PATCH 06/19] refactor: replace LongConn with ClientConn interface and simplify message handling --- internal/msggateway/client.go | 2 ++ internal/msggateway/context.go | 11 +++++++++++ 2 files changed, 13 insertions(+) diff --git a/internal/msggateway/client.go b/internal/msggateway/client.go index 46da524d3..74b874e95 100644 --- a/internal/msggateway/client.go +++ b/internal/msggateway/client.go @@ -68,6 +68,7 @@ type Client struct { UserID string `json:"userID"` IsBackground bool `json:"isBackground"` SDKType string `json:"sdkType"` + SDKVersion string `json:"sdkVersion"` Encoder Encoder ctx *UserConnContext longConnServer LongConnServer @@ -95,6 +96,7 @@ func (c *Client) ResetClient(ctx *UserConnContext, conn ClientConn, longConnServ c.closedErr = nil c.token = ctx.GetToken() c.SDKType = ctx.GetSDKType() + c.SDKVersion = ctx.GetSDKVersion() c.hbCtx, c.hbCancel = context.WithCancel(c.ctx) c.subLock = new(sync.Mutex) if c.subUserIDs != nil { diff --git a/internal/msggateway/context.go b/internal/msggateway/context.go index 6883c22a3..9fa28e667 100644 --- a/internal/msggateway/context.go +++ b/internal/msggateway/context.go @@ -38,6 +38,7 @@ type UserConnContextInfo struct { SDKType string `json:"sdkType"` SendResponse bool `json:"sendResponse"` Background bool `json:"background"` + SDKVersion string `json:"sdkVersion"` } type UserConnContext struct { @@ -74,6 +75,8 @@ func (c *UserConnContext) Value(key any) any { return c.GetPlatformID() case constant.RemoteAddr: return c.RemoteAddr + case SDKVersion: + return c.info.SDKVersion default: return "" } @@ -117,6 +120,7 @@ func (c *UserConnContext) parseByQuery(query url.Values, header http.Header) err OperationID: query.Get(OperationID), Compression: query.Get(Compression), SDKType: query.Get(SDKType), + SDKVersion: query.Get(SDKVersion), } platformID, err := strconv.Atoi(query.Get(PlatformID)) if err != nil { @@ -246,6 +250,13 @@ func (c *UserConnContext) GetSDKType() string { } } +func (c *UserConnContext) GetSDKVersion() string { + if c == nil || c.info == nil { + return "" + } + return c.info.SDKVersion +} + func (c *UserConnContext) ShouldSendResp() bool { return c != nil && c.info != nil && c.info.SendResponse } From c27d33160f167a7d967f9d1cb74a6b48c5e44ca9 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Thu, 15 Jan 2026 10:54:38 +0800 Subject: [PATCH 07/19] fix: seq use $setOnInsert for min_seq in conversation update --- pkg/common/storage/database/mgo/seq_conversation.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pkg/common/storage/database/mgo/seq_conversation.go b/pkg/common/storage/database/mgo/seq_conversation.go index 7971b7e1a..b9eb6c2a9 100644 --- a/pkg/common/storage/database/mgo/seq_conversation.go +++ b/pkg/common/storage/database/mgo/seq_conversation.go @@ -57,8 +57,8 @@ func (s *seqConversationMongo) Malloc(ctx context.Context, conversationID string } filter := map[string]any{"conversation_id": conversationID} update := map[string]any{ - "$inc": map[string]any{"max_seq": size}, - "$set": map[string]any{"min_seq": int64(0)}, + "$inc": map[string]any{"max_seq": size}, + "$setOnInsert": map[string]any{"min_seq": int64(0)}, } opt := options.FindOneAndUpdate().SetUpsert(true).SetReturnDocument(options.After).SetProjection(map[string]any{"_id": 0, "max_seq": 1}) lastSeq, err := mongoutil.FindOneAndUpdate[int64](ctx, s.coll, filter, update, opt) From 5d451fa3f86d48eaf5aab0d961bc99b265c058ec Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Thu, 22 Jan 2026 15:04:37 +0800 Subject: [PATCH 08/19] feat: add error code for handled friend requests and improve error handling in friend operations --- pkg/common/servererrs/code.go | 1 + pkg/common/servererrs/predefine.go | 9 +++++---- pkg/common/storage/controller/friend.go | 18 ++++++++---------- 3 files changed, 14 insertions(+), 14 deletions(-) diff --git a/pkg/common/servererrs/code.go b/pkg/common/servererrs/code.go index 906f890a5..153199b1b 100644 --- a/pkg/common/servererrs/code.go +++ b/pkg/common/servererrs/code.go @@ -70,6 +70,7 @@ const ( BlockedByPeer = 1302 // Blocked by the peer NotPeersFriend = 1303 // Not the peer's friend RelationshipAlreadyError = 1304 // Already in a friend relationship + FriendRequestHandled = 1305 // Friend request has already been handled // Message error codes. MessageHasReadDisable = 1401 diff --git a/pkg/common/servererrs/predefine.go b/pkg/common/servererrs/predefine.go index b1d6b06a9..06f036ae4 100644 --- a/pkg/common/servererrs/predefine.go +++ b/pkg/common/servererrs/predefine.go @@ -51,10 +51,11 @@ var ( ErrMessageHasReadDisable = errs.NewCodeError(MessageHasReadDisable, "MessageHasReadDisable") - ErrCanNotAddYourself = errs.NewCodeError(CanNotAddYourselfError, "CanNotAddYourselfError") - ErrBlockedByPeer = errs.NewCodeError(BlockedByPeer, "BlockedByPeer") - ErrNotPeersFriend = errs.NewCodeError(NotPeersFriend, "NotPeersFriend") - ErrRelationshipAlready = errs.NewCodeError(RelationshipAlreadyError, "RelationshipAlreadyError") + ErrCanNotAddYourself = errs.NewCodeError(CanNotAddYourselfError, "CanNotAddYourselfError") + ErrBlockedByPeer = errs.NewCodeError(BlockedByPeer, "BlockedByPeer") + ErrNotPeersFriend = errs.NewCodeError(NotPeersFriend, "NotPeersFriend") + ErrRelationshipAlready = errs.NewCodeError(RelationshipAlreadyError, "RelationshipAlreadyError") + ErrFriendRequestHandled = errs.NewCodeError(FriendRequestHandled, "FriendRequestHandled") ErrMutedInGroup = errs.NewCodeError(MutedInGroup, "MutedInGroup") ErrMutedGroup = errs.NewCodeError(MutedGroup, "MutedGroup") diff --git a/pkg/common/storage/controller/friend.go b/pkg/common/storage/controller/friend.go index 806468ea1..b2ae3e732 100644 --- a/pkg/common/storage/controller/friend.go +++ b/pkg/common/storage/controller/friend.go @@ -16,9 +16,9 @@ package controller import ( "context" - "fmt" "time" + "github.com/openimsdk/open-im-server/v3/pkg/common/servererrs" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database/mgo" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/model" @@ -109,15 +109,13 @@ func (f *friendDatabase) CheckIn(ctx context.Context, userID1, userID2 string) ( // Retrieve friend IDs of userID1 from the cache userID1FriendIDs, err := f.cache.GetFriendIDs(ctx, userID1) if err != nil { - err = fmt.Errorf("error retrieving friend IDs for user %s: %w", userID1, err) - return + return false, false, err } // Retrieve friend IDs of userID2 from the cache userID2FriendIDs, err := f.cache.GetFriendIDs(ctx, userID2) if err != nil { - err = fmt.Errorf("error retrieving friend IDs for user %s: %w", userID2, err) - return + return false, false, err } // Check if userID2 is in userID1's friend list and vice versa @@ -214,12 +212,12 @@ func (f *friendDatabase) RefuseFriendRequest(ctx context.Context, friendRequest // Attempt to retrieve the friend request from the database. fr, err := f.friendRequest.Take(ctx, friendRequest.FromUserID, friendRequest.ToUserID) if err != nil { - return fmt.Errorf("failed to retrieve friend request from %s to %s: %w", friendRequest.FromUserID, friendRequest.ToUserID, err) + return err } // Check if the friend request has already been handled. if fr.HandleResult != 0 { - return fmt.Errorf("friend request from %s to %s has already been processed", friendRequest.FromUserID, friendRequest.ToUserID) + return servererrs.ErrFriendRequestHandled.WrapMsg("friend request has already been processed", "from", friendRequest.FromUserID, "to", friendRequest.ToUserID) } // Log the action of refusing the friend request for debugging and auditing purposes. @@ -232,7 +230,7 @@ func (f *friendDatabase) RefuseFriendRequest(ctx context.Context, friendRequest friendRequest.HandleResult = constant.FriendResponseRefuse friendRequest.HandleTime = time.Now() if err := f.friendRequest.Update(ctx, friendRequest); err != nil { - return fmt.Errorf("failed to update friend request from %s to %s as refused: %w", friendRequest.FromUserID, friendRequest.ToUserID, err) + return err } return nil @@ -350,9 +348,9 @@ func (f *friendDatabase) PageFriendRequestToMe(ctx context.Context, userID strin func (f *friendDatabase) FindFriendsWithError(ctx context.Context, ownerUserID string, friendUserIDs []string) (friends []*model.Friend, err error) { friends, err = f.friend.FindFriends(ctx, ownerUserID, friendUserIDs) if err != nil { - return + return nil, err } - return + return friends, nil } func (f *friendDatabase) FindFriendUserIDs(ctx context.Context, ownerUserID string) (friendUserIDs []string, err error) { From 82f87551d6067983f4e31f8dacb905c740d58639 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 5 Jun 2026 10:49:47 +0800 Subject: [PATCH 09/19] refactor(msg): update regex pattern for conversationID to include a trailing colon --- pkg/common/storage/database/mgo/msg.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pkg/common/storage/database/mgo/msg.go b/pkg/common/storage/database/mgo/msg.go index 315f4530b..2d49e7c1b 100644 --- a/pkg/common/storage/database/mgo/msg.go +++ b/pkg/common/storage/database/mgo/msg.go @@ -956,7 +956,7 @@ func (m *MsgMgo) GetLastMessageSeqByTime(ctx context.Context, conversationID str { "$match": bson.M{ "doc_id": bson.M{ - "$regex": fmt.Sprintf("^%s", conversationID), + "$regex": fmt.Sprintf("^%s:", conversationID), }, }, }, @@ -1008,7 +1008,7 @@ func (m *MsgMgo) GetLastMessage(ctx context.Context, conversationID string) (*mo { "$match": bson.M{ "doc_id": bson.M{ - "$regex": fmt.Sprintf("^%s", conversationID), + "$regex": fmt.Sprintf("^%s:", conversationID), }, }, }, From d0366c483df1dbd570e89a91281bb15ffc08f5cb Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 14:33:46 +0800 Subject: [PATCH 10/19] refactor: streamline cache initialization and enhance queue engine configuration --- cmd/main.go | 80 ++++++--- internal/api/router.go | 54 +++++- internal/msggateway/init.go | 42 +---- internal/msgtransfer/init.go | 36 +--- internal/push/push.go | 30 +--- internal/rpc/auth/auth.go | 24 +-- internal/rpc/msg/server.go | 29 ++- internal/rpc/third/third.go | 27 +-- pkg/common/config/config.go | 10 +- pkg/common/config/load_config.go | 27 ++- pkg/common/config/mq.go | 41 +++++ pkg/common/config/selector.go | 43 +++++ pkg/common/discovery/discoveryregister.go | 7 +- pkg/common/storage/cache/mcache/minio.go | 50 ------ pkg/common/storage/cache/mcache/msg_cache.go | 132 -------------- pkg/common/storage/cache/mcache/online.go | 82 --------- .../storage/cache/mcache/seq_conversation.go | 79 --------- pkg/common/storage/cache/mcache/third.go | 98 ----------- pkg/common/storage/cache/mcache/token.go | 166 ------------------ pkg/common/storage/cache/mcache/tools.go | 63 ------- pkg/common/storage/cache/redis/online.go | 8 +- .../storage/cache/redis/seq_conversation.go | 7 +- pkg/mqbuild/builder.go | 143 +++++++++++++-- 23 files changed, 398 insertions(+), 880 deletions(-) create mode 100644 pkg/common/config/mq.go create mode 100644 pkg/common/config/selector.go delete mode 100644 pkg/common/storage/cache/mcache/minio.go delete mode 100644 pkg/common/storage/cache/mcache/msg_cache.go delete mode 100644 pkg/common/storage/cache/mcache/online.go delete mode 100644 pkg/common/storage/cache/mcache/seq_conversation.go delete mode 100644 pkg/common/storage/cache/mcache/third.go delete mode 100644 pkg/common/storage/cache/mcache/token.go delete mode 100644 pkg/common/storage/cache/mcache/tools.go diff --git a/cmd/main.go b/cmd/main.go index 7e19f1c98..aee83715f 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -17,7 +17,9 @@ import ( "syscall" "time" - "github.com/mitchellh/mapstructure" + "github.com/spf13/viper" + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/internal/api" "github.com/openimsdk/open-im-server/v3/internal/msggateway" "github.com/openimsdk/open-im-server/v3/internal/msgtransfer" @@ -30,33 +32,37 @@ import ( "github.com/openimsdk/open-im-server/v3/internal/rpc/third" "github.com/openimsdk/open-im-server/v3/internal/rpc/user" "github.com/openimsdk/open-im-server/v3/internal/tools/cron" + "github.com/openimsdk/open-im-server/v3/pkg/authverify" "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/open-im-server/v3/pkg/common/prommetrics" "github.com/openimsdk/open-im-server/v3/version" "github.com/openimsdk/tools/discovery" - "github.com/openimsdk/tools/discovery/standalone" + "github.com/openimsdk/tools/discovery/inprocess" "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/system/program" "github.com/openimsdk/tools/utils/datautil" - "github.com/spf13/viper" - "google.golang.org/grpc" ) func init() { config.SetStandalone() prommetrics.RegistryAll() + inprocess.SetContextAdminFunc(authverify.IsAdmin) } func main() { - var configPath string + var ( + configPath string + index int + ) flag.StringVar(&configPath, "c", "", "config path") + flag.IntVar(&index, "i", 0, "start index") flag.Parse() if configPath == "" { _, _ = fmt.Fprintln(os.Stderr, "config path is empty") os.Exit(1) return } - cmd := newCmds(configPath) + cmd := newCmds(configPath, index) putCmd(cmd, false, auth.Start) putCmd(cmd, false, conversation.Start) putCmd(cmd, false, relation.Start) @@ -77,8 +83,8 @@ func main() { } } -func newCmds(confPath string) *cmds { - return &cmds{confPath: confPath} +func newCmds(confPath string, index int) *cmds { + return &cmds{confPath: confPath, index: index} } type cmdName struct { @@ -88,6 +94,7 @@ type cmdName struct { } type cmds struct { confPath string + index int cmds []cmdName config config.AllConfig conf map[string]reflect.Value @@ -116,6 +123,9 @@ func (x *cmds) initDiscovery() { func (x *cmds) initAllConfig() error { x.conf = make(map[string]reflect.Value) + if err := x.loadShareConfig(); err != nil { + return err + } vof := reflect.ValueOf(&x.config).Elem() num := vof.NumField() for i := 0; i < num; i++ { @@ -129,23 +139,17 @@ func (x *cmds) initAllConfig() error { } x.conf[x.getTypePath(field.Type())] = field val := field.Addr().Interface() - name := val.(interface{ GetConfigFileName() string }).GetConfigFileName() - confData, err := os.ReadFile(filepath.Join(x.confPath, name)) - if err != nil { - if os.IsNotExist(err) { - continue - } - return err + // Fields without a config file (e.g. Timer) are still registered above + // for parseConf distribution, but there is nothing to load for them. + fc, ok := val.(config.FileConfig) + if !ok { + continue } - v := viper.New() - v.SetConfigType("yaml") - if err := v.ReadConfig(bytes.NewReader(confData)); err != nil { - return err + name := fc.GetConfigFileName() + if name == config.ShareFileName || x.skipKafkaConfig(name) { + continue } - opt := func(conf *mapstructure.DecoderConfig) { - conf.TagName = config.StructTagName - } - if err := v.Unmarshal(val, opt); err != nil { + if err := x.loadFileConfig(name, val); err != nil { return err } } @@ -156,6 +160,31 @@ func (x *cmds) initAllConfig() error { return nil } +func (x *cmds) loadShareConfig() error { + return x.loadFileConfig(config.ShareFileName, &x.config.Share) +} + +func (x *cmds) skipKafkaConfig(name string) bool { + return name == config.KafkaConfigFileName && + config.NormalizeQueueEngine(x.config.Share.Queue.Engine) != config.QueueEngineKafka +} + +func (x *cmds) loadFileConfig(name string, val any) error { + confData, err := os.ReadFile(filepath.Join(x.confPath, name)) + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + v := viper.New() + v.SetConfigType("yaml") + if err := v.ReadConfig(bytes.NewReader(confData)); err != nil { + return err + } + return v.Unmarshal(val, config.ApplyDecoderConfig) +} + func (x *cmds) parseConf(conf any) error { vof := reflect.ValueOf(conf) for { @@ -178,6 +207,7 @@ func (x *cmds) parseConf(conf any) error { if !ok { switch field.Interface().(type) { case config.Index: + field.Set(reflect.ValueOf(config.Index(x.index))) case config.Path: field.SetString(x.confPath) case config.AllConfig: @@ -201,7 +231,7 @@ func (x *cmds) add(name string, block bool, fn func(ctx context.Context) error) func (x *cmds) initLog() error { conf := x.config.Log if err := log.InitLoggerFromConfig( - "openim-server", + "openim-service-log", program.GetProcessName(), "", "", conf.RemainLogLevel, @@ -339,7 +369,7 @@ func putCmd[C any](cmd *cmds, block bool, fn func(ctx context.Context, config *C if err := cmd.parseConf(&conf); err != nil { return err } - return fn(ctx, &conf, standalone.GetSvcDiscoveryRegistry(), standalone.GetServiceRegistrar()) + return fn(ctx, &conf, inprocess.GetSvcDiscoveryRegistry(), inprocess.GetServiceRegistrar()) }) } diff --git a/internal/api/router.go b/internal/api/router.go index 81787f9a9..c5f93c6c0 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -9,6 +9,8 @@ import ( "github.com/gin-gonic/gin" "github.com/gin-gonic/gin/binding" "github.com/go-playground/validator/v10" + clientv3 "go.etcd.io/etcd/client/v3" + "github.com/openimsdk/open-im-server/v3/internal/api/jssdk" "github.com/openimsdk/open-im-server/v3/pkg/authverify" "github.com/openimsdk/open-im-server/v3/pkg/common/config" @@ -23,13 +25,15 @@ import ( "github.com/openimsdk/protocol/relation" "github.com/openimsdk/protocol/third" "github.com/openimsdk/protocol/user" + "github.com/openimsdk/tools/a2r" "github.com/openimsdk/tools/apiresp" "github.com/openimsdk/tools/discovery" "github.com/openimsdk/tools/discovery/etcd" + "github.com/openimsdk/tools/discovery/inprocess" "github.com/openimsdk/tools/log" + "github.com/openimsdk/tools/mcontext" "github.com/openimsdk/tools/mw" "github.com/openimsdk/tools/mw/api" - clientv3 "go.etcd.io/etcd/client/v3" ) const ( @@ -341,9 +345,55 @@ func newGinRouter(ctx context.Context, client discovery.SvcDiscoveryRegistry, cf { r.POST("/restart", cm.CheckAdmin, cm.Restart) } + + if config.Standalone() { + a := internalApi{secret: cfg.Share.Secret, client: client} + r.POST(inprocess.BroadcastPath, a.RpcInvoke) + } return r, nil } +type internalApi struct { + secret string + client discovery.SvcDiscoveryRegistry +} + +func (x *internalApi) RpcInvoke(c *gin.Context) { + req, err := a2r.ParseRequestNotCheck[inprocess.InvokeRequest](c) + if err != nil { + apiresp.GinError(c, err) + return + } + // Request length is deliberately not validated: a proto request whose fields + // are all default values marshals to zero bytes. + if req.Service == "" || req.Method == "" || req.Secret == "" { + apiresp.GinError(c, servererrs.ErrArgs) + return + } + if req.Secret != x.secret { + apiresp.GinError(c, servererrs.ErrNoPermission.WrapMsg("secret not match")) + return + } + ctx := context.Context(c) + if req.OpUserID != "" { + ctx = mcontext.SetOpUserID(ctx, req.OpUserID) + } + if req.Admin { + ctx = authverify.WithTempAdmin(ctx) + } + cc, err := x.client.GetConn(ctx, req.Service) + if err != nil { + apiresp.GinError(c, err) + return + } + var resp []byte + if err := cc.Invoke(ctx, req.Method, req.Request, &resp); err != nil { + apiresp.GinError(c, err) + return + } + apiresp.GinSuccess(c, resp) +} + func GinParseToken(authClient *rpcli.AuthClient) gin.HandlerFunc { return func(c *gin.Context) { switch c.Request.Method { @@ -385,4 +435,6 @@ func setGinIsAdmin(imAdminUserID []string) gin.HandlerFunc { var Whitelist = []string{ "/auth/get_admin_token", "/auth/parse_token", + // cross-instance rpc invoke authenticates with share.secret instead of token + inprocess.BroadcastPath, } diff --git a/internal/msggateway/init.go b/internal/msggateway/init.go index 40a57b1da..59b33ad85 100644 --- a/internal/msggateway/init.go +++ b/internal/msggateway/init.go @@ -18,13 +18,14 @@ import ( "context" "time" + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" "github.com/openimsdk/open-im-server/v3/pkg/rpccache" "github.com/openimsdk/tools/discovery" "github.com/openimsdk/tools/utils/datautil" "github.com/openimsdk/tools/utils/runtimeenv" - "google.golang.org/grpc" "github.com/openimsdk/tools/log" ) @@ -76,42 +77,3 @@ func Start(ctx context.Context, conf *Config, client discovery.SvcDiscoveryRegis return hubServer.LongConnServer.Run(ctx) } - -// -//// Start run ws server. -//func Start(ctx context.Context, index int, conf *Config) error { -// log.CInfo(ctx, "MSG-GATEWAY server is initializing", "runtimeEnv", runtimeenv.RuntimeEnvironment(), -// "rpcPorts", conf.MsgGateway.RPC.Ports, -// "wsPort", conf.MsgGateway.LongConnSvr.Ports, "prometheusPorts", conf.MsgGateway.Prometheus.Ports) -// wsPort, err := datautil.GetElemByIndex(conf.MsgGateway.LongConnSvr.Ports, index) -// if err != nil { -// return err -// } -// -// rdb, err := redisutil.NewRedisClient(ctx, conf.RedisConfig.Build()) -// if err != nil { -// return err -// } -// longServer := NewWsServer( -// conf, -// WithPort(wsPort), -// WithMaxConnNum(int64(conf.MsgGateway.LongConnSvr.WebsocketMaxConnNum)), -// WithHandshakeTimeout(time.Duration(conf.MsgGateway.LongConnSvr.WebsocketTimeout)*time.Second), -// WithMessageMaxMsgLength(conf.MsgGateway.LongConnSvr.WebsocketMaxMsgLen), -// ) -// -// hubServer := NewServer(longServer, conf, func(srv *Server) error { -// var err error -// longServer.online, err = rpccache.NewOnlineCache(srv.userClient, nil, rdb, false, longServer.subscriberUserOnlineStatusChanges) -// return err -// }) -// -// go longServer.ChangeOnlineStatus(4) -// -// netDone := make(chan error) -// go func() { -// err = hubServer.Start(ctx, index, conf) -// netDone <- err -// }() -// return hubServer.LongConnServer.Run(netDone) -//} diff --git a/internal/msgtransfer/init.go b/internal/msgtransfer/init.go index 35026c79a..bb392ff97 100644 --- a/internal/msgtransfer/init.go +++ b/internal/msgtransfer/init.go @@ -18,8 +18,6 @@ import ( "context" "fmt" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/redis" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database/mgo" "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" @@ -28,10 +26,11 @@ import ( "github.com/openimsdk/tools/mq" "github.com/openimsdk/tools/utils/runtimeenv" + "google.golang.org/grpc" + conf "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/controller" "github.com/openimsdk/tools/log" - "google.golang.org/grpc" ) type MsgTransfer struct { @@ -59,8 +58,6 @@ type Config struct { } func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryRegistry, server grpc.ServiceRegistrar) error { - builder := mqbuild.NewBuilder(&config.KafkaConfig) - log.CInfo(ctx, "MSG-TRANSFER server is initializing", "runTimeEnv", runtimeenv.RuntimeEnvironment(), "prometheusPorts", config.MsgTransfer.Prometheus.Ports, "index", config.Index) dbb := dbbuild.NewBuilder(&config.MongodbConfig, &config.RedisConfig) @@ -72,20 +69,10 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - - //if config.Discovery.Enable == conf.ETCD { - // cm := disetcd.NewConfigManager(client.(*etcd.SvcDiscoveryRegistryImpl).GetClient(), []string{ - // config.MsgTransfer.GetConfigFileName(), - // config.RedisConfig.GetConfigFileName(), - // config.MongodbConfig.GetConfigFileName(), - // config.KafkaConfig.GetConfigFileName(), - // config.Share.GetConfigFileName(), - // config.WebhooksConfig.GetConfigFileName(), - // config.Discovery.GetConfigFileName(), - // conf.LogConfigFileName, - // }) - // cm.Watch(ctx) - //} + builder, err := mqbuild.NewBuilder(config.Share.Queue, &config.KafkaConfig, rdb) + if err != nil { + return err + } mongoProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToMongoTopic) if err != nil { return err @@ -98,16 +85,7 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - var msgModel cache.MsgCache - if rdb == nil { - cm, err := mgo.NewCacheMgo(mgocli.GetDB()) - if err != nil { - return err - } - msgModel = mcache.NewMsgCache(cm, msgDocModel) - } else { - msgModel = redis.NewMsgCache(rdb, msgDocModel) - } + msgModel := redis.NewMsgCache(rdb, msgDocModel) seqConversation, err := mgo.NewSeqConversationMongo(mgocli.GetDB()) if err != nil { return err diff --git a/internal/push/push.go b/internal/push/push.go index bf95b6acc..78e6a107f 100644 --- a/internal/push/push.go +++ b/internal/push/push.go @@ -2,25 +2,24 @@ package push import ( "context" - "github.com/openimsdk/tools/mq" "math/rand" "strconv" + "github.com/openimsdk/tools/mq" + + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/internal/push/offlinepush" "github.com/openimsdk/open-im-server/v3/pkg/authverify" "github.com/openimsdk/open-im-server/v3/pkg/common/config" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/redis" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/controller" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database/mgo" "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" "github.com/openimsdk/open-im-server/v3/pkg/mqbuild" pbpush "github.com/openimsdk/protocol/push" "github.com/openimsdk/tools/discovery" "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/mcontext" - "google.golang.org/grpc" ) type pushServer struct { @@ -57,26 +56,15 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - var cacheModel cache.ThirdCache - if rdb == nil { - mdb, err := dbb.Mongo(ctx) - if err != nil { - return err - } - mc, err := mgo.NewCacheMgo(mdb.GetDB()) - if err != nil { - return err - } - cacheModel = mcache.NewThirdCache(mc) - } else { - cacheModel = redis.NewThirdCache(rdb) - } + cacheModel := redis.NewThirdCache(rdb) offlinePusher, err := offlinepush.NewOfflinePusher(&config.RpcConfig, cacheModel, string(config.FcmConfigPath)) if err != nil { return err } - builder := mqbuild.NewBuilder(&config.KafkaConfig) - + builder, err := mqbuild.NewBuilder(config.Share.Queue, &config.KafkaConfig, rdb) + if err != nil { + return err + } offlinePushProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToOfflinePushTopic) if err != nil { return err diff --git a/internal/rpc/auth/auth.go b/internal/rpc/auth/auth.go index a78a714ca..c6f64deff 100644 --- a/internal/rpc/auth/auth.go +++ b/internal/rpc/auth/auth.go @@ -19,18 +19,18 @@ import ( "errors" "github.com/openimsdk/open-im-server/v3/pkg/common/convert" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database/mgo" "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" "github.com/openimsdk/open-im-server/v3/pkg/localcache" "github.com/openimsdk/open-im-server/v3/pkg/rpccache" "github.com/openimsdk/open-im-server/v3/pkg/rpcli" + "github.com/redis/go-redis/v9" + "github.com/openimsdk/open-im-server/v3/pkg/common/config" redis2 "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/redis" "github.com/openimsdk/tools/utils/datautil" - "github.com/redis/go-redis/v9" + + "google.golang.org/grpc" "github.com/openimsdk/open-im-server/v3/pkg/authverify" "github.com/openimsdk/open-im-server/v3/pkg/common/prommetrics" @@ -43,7 +43,6 @@ import ( "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/tokenverify" - "google.golang.org/grpc" ) type authServer struct { @@ -71,20 +70,7 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - var token cache.TokenModel - if rdb == nil { - mdb, err := dbb.Mongo(ctx) - if err != nil { - return err - } - mc, err := mgo.NewCacheMgo(mdb.GetDB()) - if err != nil { - return err - } - token = mcache.NewTokenCacheModel(mc, config.RpcConfig.TokenPolicy.Expire) - } else { - token = redis2.NewTokenCacheModel(rdb, &config.LocalCacheConfig, config.RpcConfig.TokenPolicy.Expire) - } + token := redis2.NewTokenCacheModel(rdb, &config.LocalCacheConfig, config.RpcConfig.TokenPolicy.Expire) userConn, err := client.GetConn(ctx, config.Discovery.RpcService.User) if err != nil { return err diff --git a/internal/rpc/msg/server.go b/internal/rpc/msg/server.go index 48101cdd7..1d88ad933 100644 --- a/internal/rpc/msg/server.go +++ b/internal/rpc/msg/server.go @@ -17,12 +17,11 @@ package msg import ( "context" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" "github.com/openimsdk/open-im-server/v3/pkg/mqbuild" "github.com/openimsdk/open-im-server/v3/pkg/rpcli" - "google.golang.org/grpc" "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/redis" @@ -80,11 +79,6 @@ func (m *msgServer) addInterceptorHandler(interceptorFunc ...MessageInterceptorF } func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryRegistry, server grpc.ServiceRegistrar) error { - builder := mqbuild.NewBuilder(&config.KafkaConfig) - redisProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToRedisTopic) - if err != nil { - return err - } dbb := dbbuild.NewBuilder(&config.MongodbConfig, &config.RedisConfig) mgocli, err := dbb.Mongo(ctx) if err != nil { @@ -94,20 +88,19 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } + builder, err := mqbuild.NewBuilder(config.Share.Queue, &config.KafkaConfig, rdb) + if err != nil { + return err + } + redisProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToRedisTopic) + if err != nil { + return err + } msgDocModel, err := mgo.NewMsgMongo(mgocli.GetDB()) if err != nil { return err } - var msgModel cache.MsgCache - if rdb == nil { - cm, err := mgo.NewCacheMgo(mgocli.GetDB()) - if err != nil { - return err - } - msgModel = mcache.NewMsgCache(cm, msgDocModel) - } else { - msgModel = redis.NewMsgCache(rdb, msgDocModel) - } + msgModel := redis.NewMsgCache(rdb, msgDocModel) seqConversation, err := mgo.NewSeqConversationMongo(mgocli.GetDB()) if err != nil { return err diff --git a/internal/rpc/third/third.go b/internal/rpc/third/third.go index cea6a8522..90a16012d 100644 --- a/internal/rpc/third/third.go +++ b/internal/rpc/third/third.go @@ -20,8 +20,6 @@ import ( "time" "github.com/openimsdk/open-im-server/v3/pkg/authverify" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" "github.com/openimsdk/open-im-server/v3/pkg/rpcli" "github.com/openimsdk/tools/s3/disable" @@ -33,6 +31,8 @@ import ( "github.com/openimsdk/tools/s3/aws" "github.com/openimsdk/tools/s3/kodo" + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/controller" "github.com/openimsdk/protocol/third" "github.com/openimsdk/tools/discovery" @@ -40,7 +40,6 @@ import ( "github.com/openimsdk/tools/s3/cos" "github.com/openimsdk/tools/s3/minio" "github.com/openimsdk/tools/s3/oss" - "google.golang.org/grpc" ) type thirdServer struct { @@ -83,30 +82,12 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - var thirdCache cache.ThirdCache - if rdb == nil { - tc, err := mgo.NewCacheMgo(mgocli.GetDB()) - if err != nil { - return err - } - thirdCache = mcache.NewThirdCache(tc) - } else { - thirdCache = redis.NewThirdCache(rdb) - } + thirdCache := redis.NewThirdCache(rdb) // Select the oss method according to the profile policy var o s3.Interface switch enable := config.RpcConfig.Object.Enable; enable { case "minio": - var minioCache minio.Cache - if rdb == nil { - mc, err := mgo.NewCacheMgo(mgocli.GetDB()) - if err != nil { - return err - } - minioCache = mcache.NewMinioCache(mc) - } else { - minioCache = redis.NewMinioCache(rdb) - } + minioCache := redis.NewMinioCache(rdb) o, err = minio.NewMinio(ctx, minioCache, *config.MinioConfig.Build()) case "cos": o, err = cos.NewCos(*config.RpcConfig.Object.Cos.Build()) diff --git a/pkg/common/config/config.go b/pkg/common/config/config.go index 695684d15..dbcf57f20 100644 --- a/pkg/common/config/config.go +++ b/pkg/common/config/config.go @@ -28,6 +28,13 @@ import ( "github.com/openimsdk/tools/s3/oss" ) +// FileConfig is implemented by every config struct that is backed by a config +// file; GetConfigFileName reports that file's name (e.g. "redis.yml"). It is the +// single marker used to decide whether a struct field is loaded from a file. +type FileConfig interface { + GetConfigFileName() string +} + const StructTagName = "yaml" type Path string @@ -418,7 +425,8 @@ type AfterConfig struct { } type Share struct { - Secret string `yaml:"secret"` + Secret string `yaml:"secret"` + Queue EngineSelector `yaml:"queue"` IMAdminUser struct { UserIDs []string `yaml:"userIDs"` Nicknames []string `yaml:"nicknames"` diff --git a/pkg/common/config/load_config.go b/pkg/common/config/load_config.go index 142b704e1..de493ceeb 100644 --- a/pkg/common/config/load_config.go +++ b/pkg/common/config/load_config.go @@ -3,12 +3,14 @@ package config import ( "os" "path/filepath" + "reflect" "strings" "github.com/mitchellh/mapstructure" + "github.com/spf13/viper" + "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/utils/runtimeenv" - "github.com/spf13/viper" ) func Load(configDirectory string, configFileName string, envPrefix string, config any) error { @@ -35,10 +37,27 @@ func loadConfig(path string, envPrefix string, config any) error { return errs.WrapMsg(err, "failed to read config file", "path", path, "envPrefix", envPrefix) } - if err := v.Unmarshal(config, func(config *mapstructure.DecoderConfig) { - config.TagName = StructTagName - }); err != nil { + if err := v.Unmarshal(config, ApplyDecoderConfig); err != nil { return errs.WrapMsg(err, "failed to unmarshal config", "path", path, "envPrefix", envPrefix) } return nil } + +func ApplyDecoderConfig(config *mapstructure.DecoderConfig) { + config.TagName = StructTagName + config.DecodeHook = mapstructure.ComposeDecodeHookFunc( + mapstructure.StringToTimeDurationHookFunc(), + mapstructure.StringToSliceHookFunc(","), + stringToEngineSelectorHookFunc(), + ) +} + +func stringToEngineSelectorHookFunc() mapstructure.DecodeHookFuncType { + engineSelectorType := reflect.TypeOf(EngineSelector{}) + return func(from reflect.Type, to reflect.Type, data any) (any, error) { + if to != engineSelectorType || from.Kind() != reflect.String { + return data, nil + } + return EngineSelector{Engine: data.(string)}, nil + } +} diff --git a/pkg/common/config/mq.go b/pkg/common/config/mq.go new file mode 100644 index 000000000..2cba9dd38 --- /dev/null +++ b/pkg/common/config/mq.go @@ -0,0 +1,41 @@ +package config + +import ( + "strings" + + "github.com/openimsdk/tools/errs" +) + +const ( + QueueEngineKafka = "kafka" + QueueEngineRedis = "redis" + QueueEngineMemory = "memory" +) + +func NormalizeQueueEngine(engine string) string { + switch strings.ToLower(strings.TrimSpace(engine)) { + case "kafka": + return QueueEngineKafka + case "redis": + return QueueEngineRedis + case "memory": + return QueueEngineMemory + default: + return strings.ToLower(strings.TrimSpace(engine)) + } +} + +func ValidateQueueEngine(engine string, standalone bool) (string, error) { + normalized := NormalizeQueueEngine(engine) + switch normalized { + case QueueEngineKafka, QueueEngineRedis: + return normalized, nil + case QueueEngineMemory: + if standalone { + return normalized, nil + } + return "", errs.ErrArgs.WrapMsg("unsupported queue engine for microservice deployment", "engine", engine) + default: + return "", errs.ErrArgs.WrapMsg("unsupported queue engine", "engine", engine) + } +} diff --git a/pkg/common/config/selector.go b/pkg/common/config/selector.go new file mode 100644 index 000000000..8aeeeaaa3 --- /dev/null +++ b/pkg/common/config/selector.go @@ -0,0 +1,43 @@ +package config + +import "encoding/json" + +type EngineSelector struct { + Engine string `yaml:"engine" mapstructure:"engine" json:"engine"` +} + +func (e EngineSelector) String() string { + return e.Engine +} + +func (e *EngineSelector) UnmarshalYAML(unmarshal func(any) error) error { + var engine string + if err := unmarshal(&engine); err == nil { + e.Engine = engine + return nil + } + var cfg struct { + Engine string `yaml:"engine"` + } + if err := unmarshal(&cfg); err != nil { + return err + } + e.Engine = cfg.Engine + return nil +} + +func (e *EngineSelector) UnmarshalJSON(data []byte) error { + var engine string + if err := json.Unmarshal(data, &engine); err == nil { + e.Engine = engine + return nil + } + var cfg struct { + Engine string `json:"engine"` + } + if err := json.Unmarshal(data, &cfg); err != nil { + return err + } + e.Engine = cfg.Engine + return nil +} diff --git a/pkg/common/discovery/discoveryregister.go b/pkg/common/discovery/discoveryregister.go index 87333fcac..29ed5929f 100644 --- a/pkg/common/discovery/discoveryregister.go +++ b/pkg/common/discovery/discoveryregister.go @@ -17,11 +17,12 @@ package discovery import ( "time" + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/tools/discovery" - "github.com/openimsdk/tools/discovery/standalone" + "github.com/openimsdk/tools/discovery/inprocess" "github.com/openimsdk/tools/utils/runtimeenv" - "google.golang.org/grpc" "github.com/openimsdk/tools/discovery/kubernetes" @@ -32,7 +33,7 @@ import ( // NewDiscoveryRegister creates a new service discovery and registry client based on the provided environment type. func NewDiscoveryRegister(discovery *config.Discovery, watchNames []string) (discovery.SvcDiscoveryRegistry, error) { if config.Standalone() { - return standalone.GetSvcDiscoveryRegistry(), nil + return inprocess.GetSvcDiscoveryRegistry(), nil } if runtimeenv.RuntimeEnvironment() == config.KUBERNETES { return kubernetes.NewConnManager(discovery.Kubernetes.Namespace, nil, diff --git a/pkg/common/storage/cache/mcache/minio.go b/pkg/common/storage/cache/mcache/minio.go deleted file mode 100644 index f07203cc2..000000000 --- a/pkg/common/storage/cache/mcache/minio.go +++ /dev/null @@ -1,50 +0,0 @@ -package mcache - -import ( - "context" - "time" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" - "github.com/openimsdk/tools/s3/minio" -) - -func NewMinioCache(cache database.Cache) minio.Cache { - return &minioCache{ - cache: cache, - expireTime: time.Hour * 24 * 7, - } -} - -type minioCache struct { - cache database.Cache - expireTime time.Duration -} - -func (g *minioCache) getObjectImageInfoKey(key string) string { - return cachekey.GetObjectImageInfoKey(key) -} - -func (g *minioCache) getMinioImageThumbnailKey(key string, format string, width int, height int) string { - return cachekey.GetMinioImageThumbnailKey(key, format, width, height) -} - -func (g *minioCache) DelObjectImageInfoKey(ctx context.Context, keys ...string) error { - ks := make([]string, 0, len(keys)) - for _, key := range keys { - ks = append(ks, g.getObjectImageInfoKey(key)) - } - return g.cache.Del(ctx, ks) -} - -func (g *minioCache) DelImageThumbnailKey(ctx context.Context, key string, format string, width int, height int) error { - return g.cache.Del(ctx, []string{g.getMinioImageThumbnailKey(key, format, width, height)}) -} - -func (g *minioCache) GetImageObjectKeyInfo(ctx context.Context, key string, fn func(ctx context.Context) (*minio.ImageInfo, error)) (*minio.ImageInfo, error) { - return getCache[*minio.ImageInfo](ctx, g.cache, g.getObjectImageInfoKey(key), g.expireTime, fn) -} - -func (g *minioCache) GetThumbnailKey(ctx context.Context, key string, format string, width int, height int, minioCache func(ctx context.Context) (string, error)) (string, error) { - return getCache[string](ctx, g.cache, g.getMinioImageThumbnailKey(key, format, width, height), g.expireTime, minioCache) -} diff --git a/pkg/common/storage/cache/mcache/msg_cache.go b/pkg/common/storage/cache/mcache/msg_cache.go deleted file mode 100644 index 6fd5f80a1..000000000 --- a/pkg/common/storage/cache/mcache/msg_cache.go +++ /dev/null @@ -1,132 +0,0 @@ -package mcache - -import ( - "context" - "strconv" - "sync" - "time" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/model" - "github.com/openimsdk/open-im-server/v3/pkg/localcache" - "github.com/openimsdk/open-im-server/v3/pkg/localcache/lru" - "github.com/openimsdk/tools/errs" - "github.com/openimsdk/tools/utils/datautil" - "github.com/redis/go-redis/v9" -) - -var ( - memMsgCache lru.LRU[string, *model.MsgInfoModel] - initMemMsgCache sync.Once -) - -func NewMsgCache(cache database.Cache, msgDocDatabase database.Msg) cache.MsgCache { - initMemMsgCache.Do(func() { - memMsgCache = lru.NewLazyLRU[string, *model.MsgInfoModel](1024*8, time.Hour, time.Second*10, localcache.EmptyTarget{}, nil) - }) - return &msgCache{ - cache: cache, - msgDocDatabase: msgDocDatabase, - memMsgCache: memMsgCache, - } -} - -type msgCache struct { - cache database.Cache - msgDocDatabase database.Msg - memMsgCache lru.LRU[string, *model.MsgInfoModel] -} - -func (x *msgCache) getSendMsgKey(id string) string { - return cachekey.GetSendMsgKey(id) -} - -func (x *msgCache) SetSendMsgStatus(ctx context.Context, id string, status int32) error { - return x.cache.Set(ctx, x.getSendMsgKey(id), strconv.Itoa(int(status)), time.Hour*24) -} - -func (x *msgCache) GetSendMsgStatus(ctx context.Context, id string) (int32, error) { - key := x.getSendMsgKey(id) - res, err := x.cache.Get(ctx, []string{key}) - if err != nil { - return 0, err - } - val, ok := res[key] - if !ok { - return 0, errs.Wrap(redis.Nil) - } - status, err := strconv.Atoi(val) - if err != nil { - return 0, errs.WrapMsg(err, "GetSendMsgStatus strconv.Atoi error", "val", val) - } - return int32(status), nil -} - -func (x *msgCache) getMsgCacheKey(conversationID string, seq int64) string { - return cachekey.GetMsgCacheKey(conversationID, seq) - -} - -func (x *msgCache) GetMessageBySeqs(ctx context.Context, conversationID string, seqs []int64) ([]*model.MsgInfoModel, error) { - if len(seqs) == 0 { - return nil, nil - } - keys := make([]string, 0, len(seqs)) - keySeq := make(map[string]int64, len(seqs)) - for _, seq := range seqs { - key := x.getMsgCacheKey(conversationID, seq) - keys = append(keys, key) - keySeq[key] = seq - } - res, err := x.memMsgCache.GetBatch(keys, func(keys []string) (map[string]*model.MsgInfoModel, error) { - findSeqs := make([]int64, 0, len(keys)) - for _, key := range keys { - seq, ok := keySeq[key] - if !ok { - continue - } - findSeqs = append(findSeqs, seq) - } - res, err := x.msgDocDatabase.FindSeqs(ctx, conversationID, seqs) - if err != nil { - return nil, err - } - kv := make(map[string]*model.MsgInfoModel) - for i := range res { - msg := res[i] - if msg == nil || msg.Msg == nil || msg.Msg.Seq <= 0 { - continue - } - key := x.getMsgCacheKey(conversationID, msg.Msg.Seq) - kv[key] = msg - } - return kv, nil - }) - if err != nil { - return nil, err - } - return datautil.Values(res), nil -} - -func (x msgCache) DelMessageBySeqs(ctx context.Context, conversationID string, seqs []int64) error { - if len(seqs) == 0 { - return nil - } - for _, seq := range seqs { - x.memMsgCache.Del(x.getMsgCacheKey(conversationID, seq)) - } - return nil -} - -func (x *msgCache) SetMessageBySeqs(ctx context.Context, conversationID string, msgs []*model.MsgInfoModel) error { - for i := range msgs { - msg := msgs[i] - if msg == nil || msg.Msg == nil || msg.Msg.Seq <= 0 { - continue - } - x.memMsgCache.Set(x.getMsgCacheKey(conversationID, msg.Msg.Seq), msg) - } - return nil -} diff --git a/pkg/common/storage/cache/mcache/online.go b/pkg/common/storage/cache/mcache/online.go deleted file mode 100644 index f018da03e..000000000 --- a/pkg/common/storage/cache/mcache/online.go +++ /dev/null @@ -1,82 +0,0 @@ -package mcache - -import ( - "context" - "sync" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" -) - -var ( - globalOnlineCache cache.OnlineCache - globalOnlineOnce sync.Once -) - -func NewOnlineCache() cache.OnlineCache { - globalOnlineOnce.Do(func() { - globalOnlineCache = &onlineCache{ - user: make(map[string]map[int32]struct{}), - } - }) - return globalOnlineCache -} - -type onlineCache struct { - lock sync.RWMutex - user map[string]map[int32]struct{} -} - -func (x *onlineCache) GetOnline(ctx context.Context, userID string) ([]int32, error) { - x.lock.RLock() - defer x.lock.RUnlock() - pSet, ok := x.user[userID] - if !ok { - return nil, nil - } - res := make([]int32, 0, len(pSet)) - for k := range pSet { - res = append(res, k) - } - return res, nil -} - -func (x *onlineCache) SetUserOnline(ctx context.Context, userID string, online, offline []int32) error { - x.lock.Lock() - defer x.lock.Unlock() - pSet, ok := x.user[userID] - if ok { - for _, p := range offline { - delete(pSet, p) - } - } - if len(online) > 0 { - if !ok { - pSet = make(map[int32]struct{}) - x.user[userID] = pSet - } - for _, p := range online { - pSet[p] = struct{}{} - } - } - if len(pSet) == 0 { - delete(x.user, userID) - } - return nil -} - -func (x *onlineCache) GetAllOnlineUsers(ctx context.Context, cursor uint64) (map[string][]int32, uint64, error) { - if cursor != 0 { - return nil, 0, nil - } - x.lock.RLock() - defer x.lock.RUnlock() - res := make(map[string][]int32) - for k, v := range x.user { - pSet := make([]int32, 0, len(v)) - for p := range v { - pSet = append(pSet, p) - } - res[k] = pSet - } - return res, 0, nil -} diff --git a/pkg/common/storage/cache/mcache/seq_conversation.go b/pkg/common/storage/cache/mcache/seq_conversation.go deleted file mode 100644 index 879b03535..000000000 --- a/pkg/common/storage/cache/mcache/seq_conversation.go +++ /dev/null @@ -1,79 +0,0 @@ -package mcache - -import ( - "context" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" -) - -func NewSeqConversationCache(sc database.SeqConversation) cache.SeqConversationCache { - return &seqConversationCache{ - sc: sc, - } -} - -type seqConversationCache struct { - sc database.SeqConversation -} - -func (x *seqConversationCache) Malloc(ctx context.Context, conversationID string, size int64) (int64, error) { - return x.sc.Malloc(ctx, conversationID, size) -} - -func (x *seqConversationCache) SetMinSeq(ctx context.Context, conversationID string, seq int64) error { - return x.sc.SetMinSeq(ctx, conversationID, seq) -} - -func (x *seqConversationCache) GetMinSeq(ctx context.Context, conversationID string) (int64, error) { - return x.sc.GetMinSeq(ctx, conversationID) -} - -func (x *seqConversationCache) GetMaxSeqs(ctx context.Context, conversationIDs []string) (map[string]int64, error) { - res := make(map[string]int64) - for _, conversationID := range conversationIDs { - seq, err := x.GetMaxSeq(ctx, conversationID) - if err != nil { - return nil, err - } - res[conversationID] = seq - } - return res, nil -} - -func (x *seqConversationCache) GetMaxSeqsWithTime(ctx context.Context, conversationIDs []string) (map[string]database.SeqTime, error) { - res := make(map[string]database.SeqTime) - for _, conversationID := range conversationIDs { - seq, err := x.GetMaxSeq(ctx, conversationID) - if err != nil { - return nil, err - } - res[conversationID] = database.SeqTime{Seq: seq} - } - return res, nil -} - -func (x *seqConversationCache) GetMaxSeq(ctx context.Context, conversationID string) (int64, error) { - return x.sc.GetMaxSeq(ctx, conversationID) -} - -func (x *seqConversationCache) GetMaxSeqWithTime(ctx context.Context, conversationID string) (database.SeqTime, error) { - seq, err := x.GetMinSeq(ctx, conversationID) - if err != nil { - return database.SeqTime{}, err - } - return database.SeqTime{Seq: seq}, nil -} - -func (x *seqConversationCache) SetMinSeqs(ctx context.Context, seqs map[string]int64) error { - for conversationID, seq := range seqs { - if err := x.sc.SetMinSeq(ctx, conversationID, seq); err != nil { - return err - } - } - return nil -} - -func (x *seqConversationCache) GetCacheMaxSeqWithTime(ctx context.Context, conversationIDs []string) (map[string]database.SeqTime, error) { - return x.GetMaxSeqsWithTime(ctx, conversationIDs) -} diff --git a/pkg/common/storage/cache/mcache/third.go b/pkg/common/storage/cache/mcache/third.go deleted file mode 100644 index 6918ae784..000000000 --- a/pkg/common/storage/cache/mcache/third.go +++ /dev/null @@ -1,98 +0,0 @@ -package mcache - -import ( - "context" - "strconv" - "time" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" - "github.com/openimsdk/tools/errs" - "github.com/redis/go-redis/v9" -) - -func NewThirdCache(cache database.Cache) cache.ThirdCache { - return &thirdCache{ - cache: cache, - } -} - -type thirdCache struct { - cache database.Cache -} - -func (c *thirdCache) getGetuiTokenKey() string { - return cachekey.GetGetuiTokenKey() -} - -func (c *thirdCache) getGetuiTaskIDKey() string { - return cachekey.GetGetuiTaskIDKey() -} - -func (c *thirdCache) getUserBadgeUnreadCountSumKey(userID string) string { - return cachekey.GetUserBadgeUnreadCountSumKey(userID) -} - -func (c *thirdCache) getFcmAccountTokenKey(account string, platformID int) string { - return cachekey.GetFcmAccountTokenKey(account, platformID) -} - -func (c *thirdCache) get(ctx context.Context, key string) (string, error) { - res, err := c.cache.Get(ctx, []string{key}) - if err != nil { - return "", err - } - if val, ok := res[key]; ok { - return val, nil - } - return "", errs.Wrap(redis.Nil) -} - -func (c *thirdCache) SetFcmToken(ctx context.Context, account string, platformID int, fcmToken string, expireTime int64) (err error) { - return errs.Wrap(c.cache.Set(ctx, c.getFcmAccountTokenKey(account, platformID), fcmToken, time.Duration(expireTime)*time.Second)) -} - -func (c *thirdCache) GetFcmToken(ctx context.Context, account string, platformID int) (string, error) { - return c.get(ctx, c.getFcmAccountTokenKey(account, platformID)) -} - -func (c *thirdCache) DelFcmToken(ctx context.Context, account string, platformID int) error { - return c.cache.Del(ctx, []string{c.getFcmAccountTokenKey(account, platformID)}) -} - -func (c *thirdCache) IncrUserBadgeUnreadCountSum(ctx context.Context, userID string) (int, error) { - return c.cache.Incr(ctx, c.getUserBadgeUnreadCountSumKey(userID), 1) -} - -func (c *thirdCache) SetUserBadgeUnreadCountSum(ctx context.Context, userID string, value int) error { - return c.cache.Set(ctx, c.getUserBadgeUnreadCountSumKey(userID), strconv.Itoa(value), 0) -} - -func (c *thirdCache) GetUserBadgeUnreadCountSum(ctx context.Context, userID string) (int, error) { - str, err := c.get(ctx, c.getUserBadgeUnreadCountSumKey(userID)) - if err != nil { - return 0, err - } - val, err := strconv.Atoi(str) - if err != nil { - return 0, errs.WrapMsg(err, "strconv.Atoi", "str", str) - } - return val, nil -} - -func (c *thirdCache) SetGetuiToken(ctx context.Context, token string, expireTime int64) error { - return c.cache.Set(ctx, c.getGetuiTokenKey(), token, time.Duration(expireTime)*time.Second) -} - -func (c *thirdCache) GetGetuiToken(ctx context.Context) (string, error) { - return c.get(ctx, c.getGetuiTokenKey()) -} - -func (c *thirdCache) SetGetuiTaskID(ctx context.Context, taskID string, expireTime int64) error { - return c.cache.Set(ctx, c.getGetuiTaskIDKey(), taskID, time.Duration(expireTime)*time.Second) -} - -func (c *thirdCache) GetGetuiTaskID(ctx context.Context) (string, error) { - return c.get(ctx, c.getGetuiTaskIDKey()) -} diff --git a/pkg/common/storage/cache/mcache/token.go b/pkg/common/storage/cache/mcache/token.go deleted file mode 100644 index 98b9cc066..000000000 --- a/pkg/common/storage/cache/mcache/token.go +++ /dev/null @@ -1,166 +0,0 @@ -package mcache - -import ( - "context" - "fmt" - "strconv" - "strings" - "time" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" - "github.com/openimsdk/tools/errs" - "github.com/openimsdk/tools/log" -) - -func NewTokenCacheModel(cache database.Cache, accessExpire int64) cache.TokenModel { - c := &tokenCache{cache: cache} - c.accessExpire = c.getExpireTime(accessExpire) - return c -} - -type tokenCache struct { - cache database.Cache - accessExpire time.Duration -} - -func (x *tokenCache) getTokenKey(userID string, platformID int, token string) string { - return cachekey.GetTokenKey(userID, platformID) + ":" + token -} - -func (x *tokenCache) SetTokenFlag(ctx context.Context, userID string, platformID int, token string, flag int) error { - return x.cache.Set(ctx, x.getTokenKey(userID, platformID, token), strconv.Itoa(flag), x.accessExpire) -} - -// SetTokenFlagEx set token and flag with expire time -func (x *tokenCache) SetTokenFlagEx(ctx context.Context, userID string, platformID int, token string, flag int) error { - return x.SetTokenFlag(ctx, userID, platformID, token, flag) -} - -func (x *tokenCache) GetTokensWithoutError(ctx context.Context, userID string, platformID int) (map[string]int, error) { - prefix := x.getTokenKey(userID, platformID, "") - m, err := x.cache.Prefix(ctx, prefix) - if err != nil { - return nil, errs.Wrap(err) - } - mm := make(map[string]int) - for k, v := range m { - state, err := strconv.Atoi(v) - if err != nil { - log.ZError(ctx, "token value is not int", err, "value", v, "userID", userID, "platformID", platformID) - continue - } - mm[strings.TrimPrefix(k, prefix)] = state - } - return mm, nil -} - -func (x *tokenCache) HasTemporaryToken(ctx context.Context, userID string, platformID int, token string) error { - key := cachekey.GetTemporaryTokenKey(userID, platformID, token) - if _, err := x.cache.Get(ctx, []string{key}); err != nil { - return err - } - return nil -} - -func (x *tokenCache) GetAllTokensWithoutError(ctx context.Context, userID string) (map[int]map[string]int, error) { - prefix := cachekey.UidPidToken + userID + ":" - tokens, err := x.cache.Prefix(ctx, prefix) - if err != nil { - return nil, err - } - res := make(map[int]map[string]int) - for key, flagStr := range tokens { - flag, err := strconv.Atoi(flagStr) - if err != nil { - log.ZError(ctx, "token value is not int", err, "key", key, "value", flagStr, "userID", userID) - continue - } - arr := strings.SplitN(strings.TrimPrefix(key, prefix), ":", 2) - if len(arr) != 2 { - log.ZError(ctx, "token value is not int", err, "key", key, "value", flagStr, "userID", userID) - continue - } - platformID, err := strconv.Atoi(arr[0]) - if err != nil { - log.ZError(ctx, "token value is not int", err, "key", key, "value", flagStr, "userID", userID) - continue - } - token := arr[1] - if token == "" { - log.ZError(ctx, "token value is not int", err, "key", key, "value", flagStr, "userID", userID) - continue - } - tk, ok := res[platformID] - if !ok { - tk = make(map[string]int) - res[platformID] = tk - } - tk[token] = flag - } - return res, nil -} - -func (x *tokenCache) SetTokenMapByUidPid(ctx context.Context, userID string, platformID int, m map[string]int) error { - for token, flag := range m { - err := x.SetTokenFlag(ctx, userID, platformID, token, flag) - if err != nil { - return err - } - } - return nil -} - -func (x *tokenCache) BatchSetTokenMapByUidPid(ctx context.Context, tokens map[string]map[string]any) error { - for prefix, tokenFlag := range tokens { - for token, flag := range tokenFlag { - flagStr := fmt.Sprintf("%v", flag) - if err := x.cache.Set(ctx, prefix+":"+token, flagStr, x.accessExpire); err != nil { - return err - } - } - } - return nil -} - -func (x *tokenCache) DeleteTokenByUidPid(ctx context.Context, userID string, platformID int, fields []string) error { - keys := make([]string, 0, len(fields)) - for _, token := range fields { - keys = append(keys, x.getTokenKey(userID, platformID, token)) - } - return x.cache.Del(ctx, keys) -} - -func (x *tokenCache) getExpireTime(t int64) time.Duration { - return time.Hour * 24 * time.Duration(t) -} - -func (x *tokenCache) DeleteTokenByTokenMap(ctx context.Context, userID string, tokens map[int][]string) error { - keys := make([]string, 0, len(tokens)) - for platformID, ts := range tokens { - for _, t := range ts { - keys = append(keys, x.getTokenKey(userID, platformID, t)) - } - } - return x.cache.Del(ctx, keys) -} - -func (x *tokenCache) DeleteAndSetTemporary(ctx context.Context, userID string, platformID int, fields []string) error { - keys := make([]string, 0, len(fields)) - for _, f := range fields { - keys = append(keys, x.getTokenKey(userID, platformID, f)) - } - if err := x.cache.Del(ctx, keys); err != nil { - return err - } - - for _, f := range fields { - k := cachekey.GetTemporaryTokenKey(userID, platformID, f) - if err := x.cache.Set(ctx, k, "", time.Minute*5); err != nil { - return errs.Wrap(err) - } - } - - return nil -} diff --git a/pkg/common/storage/cache/mcache/tools.go b/pkg/common/storage/cache/mcache/tools.go deleted file mode 100644 index f3c4265cd..000000000 --- a/pkg/common/storage/cache/mcache/tools.go +++ /dev/null @@ -1,63 +0,0 @@ -package mcache - -import ( - "context" - "encoding/json" - "time" - - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" - "github.com/openimsdk/tools/log" -) - -func getCache[V any](ctx context.Context, cache database.Cache, key string, expireTime time.Duration, fn func(ctx context.Context) (V, error)) (V, error) { - getDB := func() (V, bool, error) { - res, err := cache.Get(ctx, []string{key}) - if err != nil { - var val V - return val, false, err - } - var val V - if str, ok := res[key]; ok { - if json.Unmarshal([]byte(str), &val) != nil { - return val, false, err - } - return val, true, nil - } - return val, false, nil - } - dbVal, ok, err := getDB() - if err != nil { - return dbVal, err - } - if ok { - return dbVal, nil - } - lockValue, err := cache.Lock(ctx, key, time.Minute) - if err != nil { - return dbVal, err - } - defer func() { - if err := cache.Unlock(ctx, key, lockValue); err != nil { - log.ZError(ctx, "unlock cache key", err, "key", key, "value", lockValue) - } - }() - dbVal, ok, err = getDB() - if err != nil { - return dbVal, err - } - if ok { - return dbVal, nil - } - val, err := fn(ctx) - if err != nil { - return val, err - } - data, err := json.Marshal(val) - if err != nil { - return val, err - } - if err := cache.Set(ctx, key, string(data), expireTime); err != nil { - return val, err - } - return val, nil -} diff --git a/pkg/common/storage/cache/redis/online.go b/pkg/common/storage/cache/redis/online.go index d09c44e9a..fe5d7ecb9 100644 --- a/pkg/common/storage/cache/redis/online.go +++ b/pkg/common/storage/cache/redis/online.go @@ -7,20 +7,16 @@ import ( "strings" "time" - "github.com/openimsdk/open-im-server/v3/pkg/common/config" + "github.com/redis/go-redis/v9" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" "github.com/openimsdk/protocol/constant" "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/log" - "github.com/redis/go-redis/v9" ) func NewUserOnline(rdb redis.UniversalClient) cache.OnlineCache { - if rdb == nil || config.Standalone() { - return mcache.NewOnlineCache() - } return &userOnline{ rdb: rdb, expire: cachekey.OnlineExpire, diff --git a/pkg/common/storage/cache/redis/seq_conversation.go b/pkg/common/storage/cache/redis/seq_conversation.go index 604826598..524f07372 100644 --- a/pkg/common/storage/cache/redis/seq_conversation.go +++ b/pkg/common/storage/cache/redis/seq_conversation.go @@ -7,20 +7,17 @@ import ( "strconv" "time" + "github.com/redis/go-redis/v9" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" - "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/mcache" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" "github.com/openimsdk/open-im-server/v3/pkg/msgprocessor" "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/log" - "github.com/redis/go-redis/v9" ) func NewSeqConversationCacheRedis(rdb redis.UniversalClient, mgo database.SeqConversation) cache.SeqConversationCache { - if rdb == nil { - return mcache.NewSeqConversationCache(mgo) - } return &seqConversationCacheRedis{ mgo: mgo, lockTime: time.Second * 3, diff --git a/pkg/mqbuild/builder.go b/pkg/mqbuild/builder.go index 938159372..c2f2d3b28 100644 --- a/pkg/mqbuild/builder.go +++ b/pkg/mqbuild/builder.go @@ -4,9 +4,12 @@ import ( "context" "fmt" + "github.com/redis/go-redis/v9" + "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/tools/mq" "github.com/openimsdk/tools/mq/kafka" + "github.com/openimsdk/tools/mq/redismq" "github.com/openimsdk/tools/mq/simmq" ) @@ -15,19 +18,120 @@ type Builder interface { GetTopicConsumer(ctx context.Context, topic string) (mq.Consumer, error) } -func NewBuilder(kafka *config.Kafka) Builder { - if config.Standalone() { - return standaloneBuilder{} +const ( + TopicToRedis = "toRedis" + TopicToMongo = "toMongo" + TopicToPush = "toPush" + TopicToOfflinePush = "toOfflinePush" +) + +const ( + GroupRedis = "redis" + GroupMongo = "mongo" + GroupPush = "push" + GroupOfflinePush = "offlinePush" +) + +func NewBuilder(queue config.EngineSelector, kafka *config.Kafka, redis redis.UniversalClient) (Builder, error) { + engine, err := config.ValidateQueueEngine(queue.Engine, config.Standalone()) + if err != nil { + return nil, err } + switch engine { + case config.QueueEngineKafka: + if kafka == nil { + return nil, fmt.Errorf("nil kafka config") + } + return newKafkaBuilder(kafka), nil + case config.QueueEngineRedis: + return newRedisBuilder(redis), nil + case config.QueueEngineMemory: + return standaloneBuilder{}, nil + default: + return nil, fmt.Errorf("unsupported queue engine %s", queue.Engine) + } +} + +func newKafkaBuilder(kafka *config.Kafka) Builder { + topics := MergeTopics(KafkaTopics(kafka)) return &kafkaBuilder{ - addr: kafka.Address, - config: kafka.Build(), - topicGroupID: map[string]string{ - kafka.ToRedisTopic: kafka.ToRedisGroupID, - kafka.ToMongoTopic: kafka.ToMongoGroupID, - kafka.ToPushTopic: kafka.ToPushGroupID, - kafka.ToOfflinePushTopic: kafka.ToOfflineGroupID, - }, + addr: kafka.Address, + config: kafka.Build(), + logicalTopic: LogicalTopicNames(topics), + topicGroupID: TopicGroupID(topics), + } +} + +func newRedisBuilder(redis redis.UniversalClient) Builder { + return redismq.NewBuilder(redis, TopicGroupID(DefaultTopics()), redismq.Config{StreamPrefix: "mq:"}) +} + +type QueueTopicsConfig struct { + ToRedis QueueTopicConfig + ToMongo QueueTopicConfig + ToPush QueueTopicConfig + ToOfflinePush QueueTopicConfig +} + +type QueueTopicConfig struct { + Topic string + GroupID string +} + +func DefaultTopics() QueueTopicsConfig { + return QueueTopicsConfig{ + ToRedis: QueueTopicConfig{Topic: TopicToRedis, GroupID: GroupRedis}, + ToMongo: QueueTopicConfig{Topic: TopicToMongo, GroupID: GroupMongo}, + ToPush: QueueTopicConfig{Topic: TopicToPush, GroupID: GroupPush}, + ToOfflinePush: QueueTopicConfig{Topic: TopicToOfflinePush, GroupID: GroupOfflinePush}, + } +} + +func KafkaTopics(kafka *config.Kafka) QueueTopicsConfig { + if kafka == nil { + return QueueTopicsConfig{} + } + return QueueTopicsConfig{ + ToRedis: QueueTopicConfig{Topic: kafka.ToRedisTopic, GroupID: kafka.ToRedisGroupID}, + ToMongo: QueueTopicConfig{Topic: kafka.ToMongoTopic, GroupID: kafka.ToMongoGroupID}, + ToPush: QueueTopicConfig{Topic: kafka.ToPushTopic, GroupID: kafka.ToPushGroupID}, + ToOfflinePush: QueueTopicConfig{Topic: kafka.ToOfflinePushTopic, GroupID: kafka.ToOfflineGroupID}, + } +} + +func MergeTopics(topics QueueTopicsConfig) QueueTopicsConfig { + defaults := DefaultTopics() + fillTopic(&topics.ToRedis, defaults.ToRedis) + fillTopic(&topics.ToMongo, defaults.ToMongo) + fillTopic(&topics.ToPush, defaults.ToPush) + fillTopic(&topics.ToOfflinePush, defaults.ToOfflinePush) + return topics +} + +func fillTopic(topic *QueueTopicConfig, defaultTopic QueueTopicConfig) { + if topic.Topic == "" { + topic.Topic = defaultTopic.Topic + } + if topic.GroupID == "" { + topic.GroupID = defaultTopic.GroupID + } +} + +func LogicalTopicNames(topics QueueTopicsConfig) map[string]string { + return map[string]string{ + TopicToRedis: topics.ToRedis.Topic, + TopicToMongo: topics.ToMongo.Topic, + TopicToPush: topics.ToPush.Topic, + TopicToOfflinePush: topics.ToOfflinePush.Topic, + } +} + +func TopicGroupID(topics QueueTopicsConfig) map[string]string { + return map[string]string{ + topics.ToRedis.Topic: topics.ToRedis.GroupID, + topics.ToMongo.Topic: topics.ToMongo.GroupID, + topics.ToPush.Topic: topics.ToPush.GroupID, + topics.ToOfflinePush.Topic: topics.ToOfflinePush.GroupID, } } @@ -44,17 +148,26 @@ func (standaloneBuilder) GetTopicConsumer(ctx context.Context, topic string) (mq type kafkaBuilder struct { addr []string config *kafka.Config + logicalTopic map[string]string topicGroupID map[string]string } func (x *kafkaBuilder) GetTopicProducer(ctx context.Context, topic string) (mq.Producer, error) { - return kafka.NewKafkaProducerV2(x.config, x.addr, topic) + realTopic, ok := x.logicalTopic[topic] + if !ok { + return nil, fmt.Errorf("topic %s not found", topic) + } + return kafka.NewKafkaProducerV2(x.config, x.addr, realTopic) } func (x *kafkaBuilder) GetTopicConsumer(ctx context.Context, topic string) (mq.Consumer, error) { - groupID, ok := x.topicGroupID[topic] + realTopic, ok := x.logicalTopic[topic] if !ok { - return nil, fmt.Errorf("topic %s groupID not found", topic) + return nil, fmt.Errorf("topic %s not found", topic) } - return kafka.NewMConsumerGroupV2(ctx, x.config, groupID, []string{topic}, true) + groupID, ok := x.topicGroupID[realTopic] + if !ok { + return nil, fmt.Errorf("topic %s groupID not found", realTopic) + } + return kafka.NewMConsumerGroupV2(ctx, x.config, groupID, []string{realTopic}, true) } From d172cfe29dc8dac0dc1dd8877f4de8c9f7987f3d Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 14:39:36 +0800 Subject: [PATCH 11/19] refactor: streamline cache initialization and enhance queue engine configuration --- cmd/main.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cmd/main.go b/cmd/main.go index aee83715f..4eed07cae 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -178,7 +178,7 @@ func (x *cmds) loadFileConfig(name string, val any) error { return err } v := viper.New() - v.SetConfigType("yaml") + v.SetConfigType(config.StructTagName) if err := v.ReadConfig(bytes.NewReader(confData)); err != nil { return err } From ad2735a1cb7b3368450d5c3eb60cbd8db81a5b08 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 15:04:40 +0800 Subject: [PATCH 12/19] feat: add RegisterIP to API config and implement Redis server registration --- cmd/main.go | 37 +++++++++++++++++++++++++++++++++++++ config/openim-api.yml | 2 ++ internal/api/init.go | 15 ++------------- pkg/common/config/config.go | 1 + 4 files changed, 42 insertions(+), 13 deletions(-) diff --git a/cmd/main.go b/cmd/main.go index 4eed07cae..e43636398 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -3,6 +3,7 @@ package main import ( "bytes" "context" + "errors" "flag" "fmt" "net" @@ -12,6 +13,7 @@ import ( "path/filepath" "reflect" "runtime" + "strconv" "strings" "sync" "syscall" @@ -41,6 +43,7 @@ import ( "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/system/program" "github.com/openimsdk/tools/utils/datautil" + "github.com/openimsdk/tools/utils/network" ) func init() { @@ -75,6 +78,7 @@ func main() { putCmd(cmd, true, msgtransfer.Start) putCmd(cmd, true, api.Start) putCmd(cmd, true, cron.Start) + putCmd(cmd, true, startRedisServerRegister) ctx := context.Background() if err := cmd.run(ctx); err != nil { _, _ = fmt.Fprintf(os.Stderr, "server exit %s", err) @@ -434,3 +438,36 @@ func (x *cmdManger) Running() []string { } return names } + +type serverConfig struct { + API config.API + Share config.Share + RedisConfig config.Redis + Index config.Index +} + +func startRedisServerRegister(ctx context.Context, cfg *serverConfig, client discovery.SvcDiscoveryRegistry, service grpc.ServiceRegistrar) error { + apiPort, err := datautil.GetElemByIndex(cfg.API.Api.Ports, int(cfg.Index)) + if err != nil { + return err + } + if apiPort <= 0 { + return errors.New("standalone api port is 0") + } + registerIP, err := network.GetRpcRegisterIP(cfg.API.Api.RegisterIP) + if err != nil { + return err + } + addr := net.JoinHostPort(registerIP, strconv.Itoa(apiPort)) + inprocess.SetLocalTarget(addr) + timer := time.NewTimer(time.Second * 5) + defer timer.Stop() + for { + select { + case <-timer.C: + case <-ctx.Done(): + return context.Cause(ctx) + } + + } +} diff --git a/config/openim-api.yml b/config/openim-api.yml index 89be50123..a879fcb1e 100644 --- a/config/openim-api.yml +++ b/config/openim-api.yml @@ -3,6 +3,8 @@ api: listenIP: 0.0.0.0 # Listening ports; if multiple are configured, multiple instances will be launched, must be consistent with the number of prometheus.ports ports: [ 10002 ] + # Fallback for IP resolution issues in multi-instance standalone mode. + # registerIP: # API compression level; 0: default compression, 1: best compression, 2: best speed, -1: no compression compressionLevel: 0 diff --git a/internal/api/init.go b/internal/api/init.go index f3548e29a..77011d0ad 100644 --- a/internal/api/init.go +++ b/internal/api/init.go @@ -23,13 +23,14 @@ import ( "strconv" "time" + "google.golang.org/grpc" + conf "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/tools/discovery" "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/utils/datautil" "github.com/openimsdk/tools/utils/network" "github.com/openimsdk/tools/utils/runtimeenv" - "google.golang.org/grpc" ) type Config struct { @@ -77,18 +78,6 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg apiCancel(err) }() - //if config.Discovery.Enable == conf.ETCD { - // cm := disetcd.NewConfigManager(client.(*etcd.SvcDiscoveryRegistryImpl).GetClient(), config.GetConfigNames()) - // cm.Watch(ctx) - //} - //sigs := make(chan os.Signal, 1) - //signal.Notify(sigs, syscall.SIGTERM) - //select { - //case val := <-sigs: - // log.ZDebug(ctx, "recv exit", "signal", val.String()) - // cancel(fmt.Errorf("signal %s", val.String())) - //case <-ctx.Done(): - //} <-apiCtx.Done() exitCause := context.Cause(apiCtx) log.ZWarn(ctx, "api server exit", exitCause) diff --git a/pkg/common/config/config.go b/pkg/common/config/config.go index dbcf57f20..8ec648359 100644 --- a/pkg/common/config/config.go +++ b/pkg/common/config/config.go @@ -142,6 +142,7 @@ type API struct { Api struct { ListenIP string `yaml:"listenIP"` Ports []int `yaml:"ports"` + RegisterIP string `yaml:"registerIP"` CompressionLevel int `yaml:"compressionLevel"` } `yaml:"api"` Prometheus struct { From ee672697c24c258c21f76d555a023df3cc8b981a Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 15:51:59 +0800 Subject: [PATCH 13/19] feat(redis): implement standalone gateway registration with Redis --- cmd/main.go | 46 ++++++++++++++-- config/share.yml | 8 +++ .../storage/cache/redis/standalone_gateway.go | 54 +++++++++++++++++++ .../cache/redis/standalone_gateway_test.go | 52 ++++++++++++++++++ 4 files changed, 155 insertions(+), 5 deletions(-) create mode 100644 pkg/common/storage/cache/redis/standalone_gateway.go create mode 100644 pkg/common/storage/cache/redis/standalone_gateway_test.go diff --git a/cmd/main.go b/cmd/main.go index e43636398..3e1c30fcc 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -37,6 +37,8 @@ import ( "github.com/openimsdk/open-im-server/v3/pkg/authverify" "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/open-im-server/v3/pkg/common/prommetrics" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/redis" + "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" "github.com/openimsdk/open-im-server/v3/version" "github.com/openimsdk/tools/discovery" "github.com/openimsdk/tools/discovery/inprocess" @@ -458,16 +460,50 @@ func startRedisServerRegister(ctx context.Context, cfg *serverConfig, client dis if err != nil { return err } - addr := net.JoinHostPort(registerIP, strconv.Itoa(apiPort)) - inprocess.SetLocalTarget(addr) - timer := time.NewTimer(time.Second * 5) - defer timer.Stop() + const validTime = time.Second * 10 + dbb := dbbuild.NewBuilder(nil, &cfg.RedisConfig) + rdb, err := dbb.Redis(ctx) + if err != nil { + return err + } + gateway := redis.NewStandaloneGatewayRedis(rdb, validTime) + selfAddr := net.JoinHostPort(registerIP, strconv.Itoa(apiPort)) + inprocess.SetLocalTarget(selfAddr) + inprocess.SetBroadcastAddress(cfg.Share.Secret, func(ctx context.Context) ([]string, error) { + address, err := gateway.GetGatewayAddrs(ctx) + if err != nil { + return nil, err + } + notSelf := make([]string, 0, len(address)) + for _, addr := range address { + if addr != selfAddr { + notSelf = append(notSelf, addr) + } + } + return notSelf, nil + }) + register := func() { + ctx, cancel := context.WithTimeout(ctx, validTime/2) + defer cancel() + if err := gateway.RegisterGateway(ctx, selfAddr); err != nil { + log.ZWarn(ctx, "gateway register failed", err, "address", selfAddr) + } + } + timer := time.NewTimer(validTime / 2) + defer func() { + timer.Stop() + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), time.Second) + defer cancel() + if err := gateway.UnregisterGateway(ctx, selfAddr); err != nil { + log.ZWarn(ctx, "gateway unregister failed", err, "address", selfAddr) + } + }() for { select { case <-timer.C: + register() case <-ctx.Done(): return context.Cause(ctx) } - } } diff --git a/config/share.yml b/config/share.yml index a42bdcdd7..df0b55be2 100644 --- a/config/share.yml +++ b/config/share.yml @@ -9,6 +9,14 @@ imAdminUser: # Each entry here corresponds by index to the matching entry in the userIDs list above. nicknames: [superAdmin] +# queue: choose message queue engine +# Supported values: +# - kafka (default) +# - redis +# - memory (standalone only; microservices cannot use memory) +queue: kafka + + # 1: For Android, iOS, Windows, Mac, and web platforms, only one instance can be online at a time multiLogin: policy: 1 diff --git a/pkg/common/storage/cache/redis/standalone_gateway.go b/pkg/common/storage/cache/redis/standalone_gateway.go new file mode 100644 index 000000000..bb46fb476 --- /dev/null +++ b/pkg/common/storage/cache/redis/standalone_gateway.go @@ -0,0 +1,54 @@ +package redis + +import ( + "context" + "strconv" + "time" + + "github.com/redis/go-redis/v9" + + "github.com/openimsdk/tools/errs" +) + +const standaloneGatewayHashKey = "STANDALONE_GATEWAY_REGISTRY" + +type StandaloneGatewayRedis struct { + rdb redis.UniversalClient + validTime time.Duration +} + +func NewStandaloneGatewayRedis(rdb redis.UniversalClient, validTime time.Duration) *StandaloneGatewayRedis { + return &StandaloneGatewayRedis{rdb: rdb, validTime: validTime} +} + +func (s *StandaloneGatewayRedis) RegisterGateway(ctx context.Context, addr string) error { + pipe := s.rdb.Pipeline() + pipe.HSet(ctx, standaloneGatewayHashKey, addr, strconv.FormatInt(time.Now().UnixMilli(), 10)) + pipe.Expire(ctx, standaloneGatewayHashKey, s.validTime*2) + _, err := pipe.Exec(ctx) + return errs.Wrap(err) +} + +func (s *StandaloneGatewayRedis) UnregisterGateway(ctx context.Context, addr string) error { + return errs.Wrap(s.rdb.HDel(ctx, standaloneGatewayHashKey, addr).Err()) +} + +func (s *StandaloneGatewayRedis) GetGatewayAddrs(ctx context.Context) ([]string, error) { + gateways, err := s.rdb.HGetAll(ctx, standaloneGatewayHashKey).Result() + if err != nil { + return nil, errs.Wrap(err) + } + + now := time.Now() + addrs := make([]string, 0, len(gateways)) + for addr, registeredAt := range gateways { + registeredAtMs, err := strconv.ParseInt(registeredAt, 10, 64) + if err != nil { + return nil, errs.WrapMsg(err, "redis gateway register time is not int64", "addr", addr, "value", registeredAt) + } + if now.Sub(time.UnixMilli(registeredAtMs)) <= s.validTime { + addrs = append(addrs, addr) + } + } + return addrs, nil +} diff --git a/pkg/common/storage/cache/redis/standalone_gateway_test.go b/pkg/common/storage/cache/redis/standalone_gateway_test.go new file mode 100644 index 000000000..b4f2c025f --- /dev/null +++ b/pkg/common/storage/cache/redis/standalone_gateway_test.go @@ -0,0 +1,52 @@ +package redis + +import ( + "context" + "strconv" + "testing" + "time" + + "github.com/go-redis/redismock/v9" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestStandaloneGatewayRedisRegisterGateway(t *testing.T) { + rdb, mock := redismock.NewClientMock() + cache := NewStandaloneGatewayRedis(rdb, time.Second*10) + + mock.Regexp().ExpectHSet(standaloneGatewayHashKey, "127.0.0.1:10001", `^[0-9]+$`).SetVal(1) + mock.ExpectExpire(standaloneGatewayHashKey, time.Second*20).SetVal(true) + + err := cache.RegisterGateway(context.Background(), "127.0.0.1:10001") + require.NoError(t, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestStandaloneGatewayRedisUnregisterGateway(t *testing.T) { + rdb, mock := redismock.NewClientMock() + cache := NewStandaloneGatewayRedis(rdb, time.Second) + + mock.ExpectHDel(standaloneGatewayHashKey, "127.0.0.1:10001").SetVal(1) + + err := cache.UnregisterGateway(context.Background(), "127.0.0.1:10001") + require.NoError(t, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestStandaloneGatewayRedisGetGatewayAddrs(t *testing.T) { + rdb, mock := redismock.NewClientMock() + cache := NewStandaloneGatewayRedis(rdb, time.Second*10) + + now := time.Now() + mock.ExpectHGetAll(standaloneGatewayHashKey).SetVal(map[string]string{ + "127.0.0.1:10001": strconv.FormatInt(now.Add(-time.Second).UnixMilli(), 10), + "127.0.0.1:10002": strconv.FormatInt(now.Add(-time.Second*20).UnixMilli(), 10), + "127.0.0.1:10003": strconv.FormatInt(now.Add(time.Second).UnixMilli(), 10), + }) + + addrs, err := cache.GetGatewayAddrs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"127.0.0.1:10001", "127.0.0.1:10003"}, addrs) + assert.NoError(t, mock.ExpectationsWereMet()) +} From 26e22efdd876965d6e55806e8560595e19db469d Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 16:19:12 +0800 Subject: [PATCH 14/19] chore: update openimsdk/tools dependency to v0.0.50-alpha.121 --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index f00a6ee40..8a8dae98f 100644 --- a/go.mod +++ b/go.mod @@ -13,7 +13,7 @@ require ( github.com/grpc-ecosystem/go-grpc-prometheus v1.2.0 github.com/mitchellh/mapstructure v1.5.0 github.com/openimsdk/protocol v0.0.73-alpha.19 - github.com/openimsdk/tools v0.0.50-alpha.117 + github.com/openimsdk/tools v0.0.50-alpha.121 github.com/pkg/errors v0.9.1 // indirect github.com/prometheus/client_golang v1.18.0 github.com/stretchr/testify v1.11.1 diff --git a/go.sum b/go.sum index 0da371b9f..5d0c9e012 100644 --- a/go.sum +++ b/go.sum @@ -363,8 +363,8 @@ github.com/openimsdk/gomake v0.0.17 h1:q8haP48VOH45WhJRiLj1YSBJyUFJqD8CTedH65i1Y github.com/openimsdk/gomake v0.0.17/go.mod h1:nnjS8yCtrPJAt1knMbyPiUwCH2gpyBzj/EZAONfUOXg= github.com/openimsdk/protocol v0.0.73-alpha.19 h1:CvXoDF2U73UcMhLnrtMFks2Aw+bXiDgH8AITEt783/s= github.com/openimsdk/protocol v0.0.73-alpha.19/go.mod h1:WF7EuE55vQvpyUAzDXcqg+B+446xQyEba0X35lTINmw= -github.com/openimsdk/tools v0.0.50-alpha.117 h1:ACfijEVCeBcttT7OOkNGOOOvq14pJtb9szNIMHLm6Vc= -github.com/openimsdk/tools v0.0.50-alpha.117/go.mod h1:I0WESSa7ghPIo9BL+ETlH/qEIbO6+KZioM1jwNuDwz0= +github.com/openimsdk/tools v0.0.50-alpha.121 h1:TXKKgtkeMeqIs0vpolbW8rIEngE9xlESq+0NV+FoLH0= +github.com/openimsdk/tools v0.0.50-alpha.121/go.mod h1:I0WESSa7ghPIo9BL+ETlH/qEIbO6+KZioM1jwNuDwz0= github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6Wq+LM= github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs= github.com/pierrec/lz4/v4 v4.1.21 h1:yOVMLb6qSIDP67pl/5F7RepeKYu/VmTyEXvuMI5d9mQ= From fa411a79a41dd1bc85987453d507e5e9e2d932e9 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 16:28:20 +0800 Subject: [PATCH 15/19] fix: change timer to ticker for improved periodic execution --- cmd/main.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cmd/main.go b/cmd/main.go index 3e1c30fcc..e5eefece5 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -489,7 +489,7 @@ func startRedisServerRegister(ctx context.Context, cfg *serverConfig, client dis log.ZWarn(ctx, "gateway register failed", err, "address", selfAddr) } } - timer := time.NewTimer(validTime / 2) + timer := time.NewTicker(validTime / 2) defer func() { timer.Stop() ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), time.Second) From bfe8bcd9ad2e455466999daa9f39329d4b28eeab Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 17:16:50 +0800 Subject: [PATCH 16/19] refactor: replace Kafka topic configuration with mqbuild constants for improved readability --- cmd/main.go | 1 + internal/msgtransfer/init.go | 8 ++++---- internal/push/push.go | 6 +++--- internal/rpc/msg/server.go | 2 +- pkg/common/config/mq.go | 2 +- 5 files changed, 10 insertions(+), 9 deletions(-) diff --git a/cmd/main.go b/cmd/main.go index e5eefece5..7b0a7a42f 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -498,6 +498,7 @@ func startRedisServerRegister(ctx context.Context, cfg *serverConfig, client dis log.ZWarn(ctx, "gateway unregister failed", err, "address", selfAddr) } }() + register() for { select { case <-timer.C: diff --git a/internal/msgtransfer/init.go b/internal/msgtransfer/init.go index bb392ff97..2b483fcd7 100644 --- a/internal/msgtransfer/init.go +++ b/internal/msgtransfer/init.go @@ -73,11 +73,11 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - mongoProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToMongoTopic) + mongoProducer, err := builder.GetTopicProducer(ctx, mqbuild.TopicToMongo) if err != nil { return err } - pushProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToPushTopic) + pushProducer, err := builder.GetTopicProducer(ctx, mqbuild.TopicToPush) if err != nil { return err } @@ -100,11 +100,11 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - historyConsumer, err := builder.GetTopicConsumer(ctx, config.KafkaConfig.ToRedisTopic) + historyConsumer, err := builder.GetTopicConsumer(ctx, mqbuild.TopicToRedis) if err != nil { return err } - historyMongoConsumer, err := builder.GetTopicConsumer(ctx, config.KafkaConfig.ToMongoTopic) + historyMongoConsumer, err := builder.GetTopicConsumer(ctx, mqbuild.TopicToMongo) if err != nil { return err } diff --git a/internal/push/push.go b/internal/push/push.go index 78e6a107f..d68708422 100644 --- a/internal/push/push.go +++ b/internal/push/push.go @@ -65,17 +65,17 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - offlinePushProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToOfflinePushTopic) + offlinePushProducer, err := builder.GetTopicProducer(ctx, mqbuild.TopicToOfflinePush) if err != nil { return err } database := controller.NewPushDatabase(cacheModel, offlinePushProducer) - pushConsumer, err := builder.GetTopicConsumer(ctx, config.KafkaConfig.ToPushTopic) + pushConsumer, err := builder.GetTopicConsumer(ctx, mqbuild.TopicToPush) if err != nil { return err } - offlinePushConsumer, err := builder.GetTopicConsumer(ctx, config.KafkaConfig.ToOfflinePushTopic) + offlinePushConsumer, err := builder.GetTopicConsumer(ctx, mqbuild.TopicToOfflinePush) if err != nil { return err } diff --git a/internal/rpc/msg/server.go b/internal/rpc/msg/server.go index 1d88ad933..e7c589b8c 100644 --- a/internal/rpc/msg/server.go +++ b/internal/rpc/msg/server.go @@ -92,7 +92,7 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg if err != nil { return err } - redisProducer, err := builder.GetTopicProducer(ctx, config.KafkaConfig.ToRedisTopic) + redisProducer, err := builder.GetTopicProducer(ctx, mqbuild.TopicToRedis) if err != nil { return err } diff --git a/pkg/common/config/mq.go b/pkg/common/config/mq.go index 2cba9dd38..76fdb26e6 100644 --- a/pkg/common/config/mq.go +++ b/pkg/common/config/mq.go @@ -14,7 +14,7 @@ const ( func NormalizeQueueEngine(engine string) string { switch strings.ToLower(strings.TrimSpace(engine)) { - case "kafka": + case "", "kafka": return QueueEngineKafka case "redis": return QueueEngineRedis From 5882597ee1a3cf31dd38a2fff08378a3b30bf195 Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Fri, 3 Jul 2026 17:51:42 +0800 Subject: [PATCH 17/19] feat(redis): implement Redis-based locking mechanism for cron tasks --- internal/tools/cron/cron_task.go | 34 ++++------ .../cron/{dist_look.go => etcd_locker.go} | 7 +- internal/tools/cron/locker.go | 14 ++++ internal/tools/cron/redis_locker.go | 68 +++++++++++++++++++ pkg/common/cmd/cron_task.go | 4 +- .../cache/redis/standalone_gateway_test.go | 2 +- 6 files changed, 103 insertions(+), 26 deletions(-) rename internal/tools/cron/{dist_look.go => etcd_locker.go} (97%) create mode 100644 internal/tools/cron/locker.go create mode 100644 internal/tools/cron/redis_locker.go diff --git a/internal/tools/cron/cron_task.go b/internal/tools/cron/cron_task.go index 2c8655d4a..e4644435b 100644 --- a/internal/tools/cron/cron_task.go +++ b/internal/tools/cron/cron_task.go @@ -3,8 +3,12 @@ package cron import ( "context" + "github.com/robfig/cron/v3" + "google.golang.org/grpc" + "github.com/openimsdk/open-im-server/v3/pkg/common/config" disetcd "github.com/openimsdk/open-im-server/v3/pkg/common/discovery/etcd" + "github.com/openimsdk/open-im-server/v3/pkg/dbbuild" pbconversation "github.com/openimsdk/protocol/conversation" "github.com/openimsdk/protocol/msg" "github.com/openimsdk/protocol/third" @@ -14,14 +18,13 @@ import ( "github.com/openimsdk/tools/log" "github.com/openimsdk/tools/mcontext" "github.com/openimsdk/tools/utils/runtimeenv" - "github.com/robfig/cron/v3" - "google.golang.org/grpc" ) type Config struct { - CronTask config.CronTask - Share config.Share - Discovery config.Discovery + CronTask config.CronTask + Share config.Share + Discovery config.Discovery + RedisConfig config.Redis } func Start(ctx context.Context, conf *Config, client discovery.SvcDiscoveryRegistry, service grpc.ServiceRegistrar) error { @@ -60,10 +63,13 @@ func Start(ctx context.Context, conf *Config, client discovery.SvcDiscoveryRegis if err != nil { return err } - } - - if locker == nil { - locker = emptyLocker{} + } else { + builder := dbbuild.NewBuilder(nil, &conf.RedisConfig) + rdb, err := builder.Redis(ctx) + if err != nil { + return err + } + locker = NewRedisLocker(rdb) } srv := &cronServer{ @@ -95,16 +101,6 @@ func Start(ctx context.Context, conf *Config, client discovery.SvcDiscoveryRegis return nil } -type Locker interface { - ExecuteWithLock(ctx context.Context, taskName string, task func()) -} - -type emptyLocker struct{} - -func (emptyLocker) ExecuteWithLock(ctx context.Context, taskName string, task func()) { - task() -} - type cronServer struct { ctx context.Context config *Config diff --git a/internal/tools/cron/dist_look.go b/internal/tools/cron/etcd_locker.go similarity index 97% rename from internal/tools/cron/dist_look.go rename to internal/tools/cron/etcd_locker.go index e46b9206c..959110fea 100644 --- a/internal/tools/cron/dist_look.go +++ b/internal/tools/cron/etcd_locker.go @@ -6,13 +6,10 @@ import ( "os" "time" - "github.com/openimsdk/tools/log" clientv3 "go.etcd.io/etcd/client/v3" "go.etcd.io/etcd/client/v3/concurrency" -) -const ( - lockLeaseTTL = 300 + "github.com/openimsdk/tools/log" ) type EtcdLocker struct { @@ -35,7 +32,7 @@ func NewEtcdLocker(client *clientv3.Client) (*EtcdLocker, error) { } func (e *EtcdLocker) ExecuteWithLock(ctx context.Context, taskName string, task func()) { - session, err := concurrency.NewSession(e.client, concurrency.WithTTL(lockLeaseTTL)) + session, err := concurrency.NewSession(e.client, concurrency.WithTTL(int(lockLeaseTTL/time.Second))) if err != nil { log.ZWarn(ctx, "Failed to create etcd session", err, "taskName", taskName, diff --git a/internal/tools/cron/locker.go b/internal/tools/cron/locker.go new file mode 100644 index 000000000..4a1115486 --- /dev/null +++ b/internal/tools/cron/locker.go @@ -0,0 +1,14 @@ +package cron + +import ( + "context" + "time" +) + +const ( + lockLeaseTTL = time.Second * 300 +) + +type Locker interface { + ExecuteWithLock(ctx context.Context, taskName string, task func()) +} diff --git a/internal/tools/cron/redis_locker.go b/internal/tools/cron/redis_locker.go new file mode 100644 index 000000000..329e9cb77 --- /dev/null +++ b/internal/tools/cron/redis_locker.go @@ -0,0 +1,68 @@ +package cron + +import ( + "context" + "strings" + "time" + + "github.com/google/uuid" + "github.com/redis/go-redis/v9" + + "github.com/openimsdk/tools/log" +) + +func NewRedisLocker(client redis.UniversalClient) *RedisLocker { + return &RedisLocker{ + client: client, + script: redis.NewScript(strings.TrimSpace(` +if redis.call("get", KEYS[1]) == ARGV[1] then + return redis.call("del", KEYS[1]) +else + return 0 +end +`)), + } +} + +type RedisLocker struct { + client redis.UniversalClient + script *redis.Script +} + +func (e *RedisLocker) getKey(name string) string { + return "CRON_LOCKED:" + name +} + +func (e *RedisLocker) lock(ctx context.Context, name string, owner string) (bool, error) { + ctx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + return e.client.SetNX(ctx, e.getKey(name), owner, lockLeaseTTL).Result() +} + +func (e *RedisLocker) unlock(ctx context.Context, name string, owner string) error { + ctx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + return e.script.Run(ctx, e.client, []string{e.getKey(name)}, owner).Err() +} + +func (e *RedisLocker) ExecuteWithLock(ctx context.Context, taskName string, task func()) { + owner := uuid.New().String() + ok, err := e.lock(ctx, taskName, owner) + if err != nil { + log.ZWarn(ctx, "cron lock get lock", err, "taskName", taskName) + return + } + log.ZDebug(ctx, "cron lock get lock", "taskName", taskName, "ok", ok, "owner", owner) + if !ok { + return + } + defer func() { + err := e.unlock(ctx, taskName, owner) + if err == nil { + log.ZDebug(ctx, "cron lock unlock", "taskName", taskName, "owner", owner) + } else { + log.ZWarn(ctx, "cron lock unlock", err, "taskName", taskName, "owner", owner) + } + }() + task() +} diff --git a/pkg/common/cmd/cron_task.go b/pkg/common/cmd/cron_task.go index c666bd021..041deea67 100644 --- a/pkg/common/cmd/cron_task.go +++ b/pkg/common/cmd/cron_task.go @@ -17,12 +17,13 @@ package cmd import ( "context" + "github.com/spf13/cobra" + "github.com/openimsdk/open-im-server/v3/internal/tools/cron" "github.com/openimsdk/open-im-server/v3/pkg/common/config" "github.com/openimsdk/open-im-server/v3/pkg/common/startrpc" "github.com/openimsdk/open-im-server/v3/version" "github.com/openimsdk/tools/system/program" - "github.com/spf13/cobra" ) type CronTaskCmd struct { @@ -39,6 +40,7 @@ func NewCronTaskCmd() *CronTaskCmd { config.OpenIMCronTaskCfgFileName: &cronTaskConfig.CronTask, config.ShareFileName: &cronTaskConfig.Share, config.DiscoveryConfigFilename: &cronTaskConfig.Discovery, + config.RedisConfigFileName: &cronTaskConfig.RedisConfig, } ret.RootCmd = NewRootCmd(program.GetProcessName(), WithConfigMap(ret.configMap)) ret.ctx = context.WithValue(context.Background(), "version", version.Version) diff --git a/pkg/common/storage/cache/redis/standalone_gateway_test.go b/pkg/common/storage/cache/redis/standalone_gateway_test.go index b4f2c025f..3b88df104 100644 --- a/pkg/common/storage/cache/redis/standalone_gateway_test.go +++ b/pkg/common/storage/cache/redis/standalone_gateway_test.go @@ -47,6 +47,6 @@ func TestStandaloneGatewayRedisGetGatewayAddrs(t *testing.T) { addrs, err := cache.GetGatewayAddrs(context.Background()) require.NoError(t, err) - assert.Equal(t, []string{"127.0.0.1:10001", "127.0.0.1:10003"}, addrs) + assert.ElementsMatch(t, []string{"127.0.0.1:10001", "127.0.0.1:10003"}, addrs) assert.NoError(t, mock.ExpectationsWereMet()) } From 10baf8ff7e187872ed2968d79a1b96aef177efdd Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Mon, 13 Jul 2026 14:24:49 +0800 Subject: [PATCH 18/19] feat: Support Streaming Messages in the Open-Source Server --- go.mod | 2 +- go.sum | 4 +- internal/api/msg.go | 7 +- internal/api/router.go | 5 +- internal/api/stream_msg.go | 216 ++++++++++++++++++ internal/rpc/msg/modify.go | 167 ++++++++++++++ internal/rpc/msg/notification.go | 4 + internal/rpc/msg/send.go | 5 + internal/rpc/msg/server.go | 5 + internal/rpc/msg/stream_msg.go | 193 ++++++++++++++++ pkg/apistruct/msg.go | 7 +- pkg/common/storage/cache/cachekey/msg.go | 5 + .../storage/cache/cachekey/stream_msg.go | 7 + pkg/common/storage/cache/lock.go | 11 + pkg/common/storage/cache/msg.go | 2 + pkg/common/storage/cache/redis/lock.go | 60 +++++ pkg/common/storage/cache/redis/msg.go | 10 + pkg/common/storage/cache/redis/stream_msg.go | 140 ++++++++++++ pkg/common/storage/cache/stream_msg.go | 21 ++ pkg/common/storage/controller/msg.go | 28 +++ pkg/common/storage/controller/stream_msg.go | 15 ++ pkg/notification/msg.go | 7 +- 22 files changed, 910 insertions(+), 11 deletions(-) create mode 100644 internal/api/stream_msg.go create mode 100644 internal/rpc/msg/modify.go create mode 100644 internal/rpc/msg/stream_msg.go create mode 100644 pkg/common/storage/cache/cachekey/stream_msg.go create mode 100644 pkg/common/storage/cache/lock.go create mode 100644 pkg/common/storage/cache/redis/lock.go create mode 100644 pkg/common/storage/cache/redis/stream_msg.go create mode 100644 pkg/common/storage/cache/stream_msg.go create mode 100644 pkg/common/storage/controller/stream_msg.go diff --git a/go.mod b/go.mod index 8a8dae98f..61349964e 100644 --- a/go.mod +++ b/go.mod @@ -12,7 +12,7 @@ require ( github.com/gorilla/websocket v1.5.1 github.com/grpc-ecosystem/go-grpc-prometheus v1.2.0 github.com/mitchellh/mapstructure v1.5.0 - github.com/openimsdk/protocol v0.0.73-alpha.19 + github.com/openimsdk/protocol v0.0.73-alpha.20 github.com/openimsdk/tools v0.0.50-alpha.121 github.com/pkg/errors v0.9.1 // indirect github.com/prometheus/client_golang v1.18.0 diff --git a/go.sum b/go.sum index 5d0c9e012..f4b070a3f 100644 --- a/go.sum +++ b/go.sum @@ -361,8 +361,8 @@ github.com/onsi/gomega v1.25.0 h1:Vw7br2PCDYijJHSfBOWhov+8cAnUf8MfMaIOV323l6Y= github.com/onsi/gomega v1.25.0/go.mod h1:r+zV744Re+DiYCIPRlYOTxn0YkOLcAnW8k1xXdMPGhM= github.com/openimsdk/gomake v0.0.17 h1:q8haP48VOH45WhJRiLj1YSBJyUFJqD8CTedH65i1YH8= github.com/openimsdk/gomake v0.0.17/go.mod h1:nnjS8yCtrPJAt1knMbyPiUwCH2gpyBzj/EZAONfUOXg= -github.com/openimsdk/protocol v0.0.73-alpha.19 h1:CvXoDF2U73UcMhLnrtMFks2Aw+bXiDgH8AITEt783/s= -github.com/openimsdk/protocol v0.0.73-alpha.19/go.mod h1:WF7EuE55vQvpyUAzDXcqg+B+446xQyEba0X35lTINmw= +github.com/openimsdk/protocol v0.0.73-alpha.20 h1:9MnACSi6IKv2iqlxHYUJG9mgt9gyPRHxE7Lq8dxAoxI= +github.com/openimsdk/protocol v0.0.73-alpha.20/go.mod h1:WF7EuE55vQvpyUAzDXcqg+B+446xQyEba0X35lTINmw= github.com/openimsdk/tools v0.0.50-alpha.121 h1:TXKKgtkeMeqIs0vpolbW8rIEngE9xlESq+0NV+FoLH0= github.com/openimsdk/tools v0.0.50-alpha.121/go.mod h1:I0WESSa7ghPIo9BL+ETlH/qEIbO6+KZioM1jwNuDwz0= github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6Wq+LM= diff --git a/internal/api/msg.go b/internal/api/msg.go index 06fd14936..f93163a65 100644 --- a/internal/api/msg.go +++ b/internal/api/msg.go @@ -79,12 +79,13 @@ func getMsgDataDescriptor() []protoreflect.FieldDescriptor { type MessageApi struct { Client msg.MsgClient userClient *rpcli.UserClient + authClient *rpcli.AuthClient imAdminUserID []string validate *validator.Validate } -func NewMessageApi(client msg.MsgClient, userClient *rpcli.UserClient, imAdminUserID []string) MessageApi { - return MessageApi{Client: client, userClient: userClient, imAdminUserID: imAdminUserID, validate: validator.New()} +func NewMessageApi(client msg.MsgClient, userClient *rpcli.UserClient, authClient *rpcli.AuthClient, imAdminUserID []string) MessageApi { + return MessageApi{Client: client, userClient: userClient, authClient: authClient, imAdminUserID: imAdminUserID, validate: validator.New()} } func (*MessageApi) SetOptions(options map[string]bool, value bool) { @@ -219,6 +220,8 @@ func (m *MessageApi) getSendMsgReq(c *gin.Context, req apistruct.SendMsg) (sendM data = &apistruct.CustomElem{} case constant.MarkdownText: data = &apistruct.MarkdownTextElem{} + case constant.Stream: + data = &apistruct.StreamMsgElem{} case constant.Quote: data = &apistruct.QuoteElem{} case constant.OANotification: diff --git a/internal/api/router.go b/internal/api/router.go index c5f93c6c0..0a7f74394 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -253,7 +253,7 @@ func newGinRouter(ctx context.Context, client discovery.SvcDiscoveryRegistry, cf objectGroup.GET("/*name", t.ObjectRedirect) } // Message - m := NewMessageApi(msg.NewMsgClient(msgConn), rpcli.NewUserClient(userConn), cfg.Share.IMAdminUser.UserIDs) + m := NewMessageApi(msg.NewMsgClient(msgConn), rpcli.NewUserClient(userConn), rpcli.NewAuthClient(authConn), cfg.Share.IMAdminUser.UserIDs) { msgGroup := r.Group("/msg") msgGroup.POST("/newest_seq", m.GetSeq) @@ -277,6 +277,9 @@ func newGinRouter(ctx context.Context, client discovery.SvcDiscoveryRegistry, cf msgGroup.POST("/send_simple_msg", m.SendSimpleMessage) msgGroup.POST("/check_msg_is_send_success", m.CheckMsgIsSendSuccess) msgGroup.POST("/get_server_time", m.GetServerTime) + msgGroup.POST("/get_stream_msg", m.GetStreamMsg) + msgGroup.POST("/append_stream_msg", m.AppendStreamMsg) + msgGroup.PUT("/append_stream_msg", m.PutStreamMsg) } // Conversation { diff --git a/internal/api/stream_msg.go b/internal/api/stream_msg.go new file mode 100644 index 000000000..5df4c4035 --- /dev/null +++ b/internal/api/stream_msg.go @@ -0,0 +1,216 @@ +package api + +import ( + "bufio" + "bytes" + "context" + "fmt" + "io" + "net/http" + "time" + "unicode/utf8" + + "github.com/gin-gonic/gin" + + "github.com/openimsdk/protocol/constant" + "github.com/openimsdk/protocol/msg" + "github.com/openimsdk/tools/a2r" + "github.com/openimsdk/tools/apiresp" + "github.com/openimsdk/tools/errs" + "github.com/openimsdk/tools/log" +) + +func (m *MessageApi) GetStreamMsg(c *gin.Context) { + a2r.Call(c, msg.MsgClient.GetStreamMsg, m.Client) +} + +func (m *MessageApi) AppendStreamMsg(c *gin.Context) { + a2r.Call(c, msg.MsgClient.AppendStreamMsg, m.Client) +} + +func (m *MessageApi) PutStreamMsg(c *gin.Context) { + var ( + conversationID string + clientMsgID string + ) + { + operationID := c.GetHeader(constant.OperationID) + if operationID == "" { + operationID = c.Query(constant.OperationID) + } + if operationID == "" { + m.putErr(c, errs.ErrArgs.WrapMsg("operationID is empty")) + return + } + c.Set(constant.OperationID, operationID) + conversationID = c.Query("conversationID") + if conversationID == "" { + conversationID = c.GetHeader("conversationID") + } + if conversationID == "" { + m.putErr(c, errs.ErrArgs.WrapMsg("conversationID is empty")) + return + } + clientMsgID = c.Query("clientMsgID") + if clientMsgID == "" { + clientMsgID = c.GetHeader("clientMsgID") + } + if clientMsgID == "" { + m.putErr(c, errs.ErrArgs.WrapMsg("clientMsgID is empty")) + return + } + token := c.GetHeader("token") + if token == "" { + token = c.Query("token") + } + if token == "" { + m.putErr(c, errs.ErrTokenInvalid.WrapMsg("token is empty")) + return + } + resp, err := m.authClient.ParseToken(c, token) + if err != nil { + m.putErr(c, err) + return + } + c.Set(constant.OpUserPlatform, constant.PlatformIDToName(int(resp.PlatformID))) + c.Set(constant.OpUserID, resp.UserID) + } + done := make(chan struct{}) + streamCh := make(chan string, 8) + + go func() { + defer func() { + close(streamCh) + c.Request.Body.Close() + }() + buf := make([]byte, 256) + body := NewUTF8Reader(c.Request.Body) + for i := 1; ; i++ { + n, err := body.Read(buf) + if n > 0 { + select { + case streamCh <- string(buf[:n]): + case <-done: + return + } + } + if err != nil { + if err == io.EOF { + log.ZDebug(c, "read request body stream msg done", "clientMsgID", clientMsgID) + } else { + log.ZError(c, "read request body stream msg failed", err, "clientMsgID", clientMsgID, "error", err) + } + return + } + if n < 10 { + time.Sleep(time.Millisecond * 10) + } + } + }() + + var ( + packet []string + end bool + index int + errCount int + lastErr error + ) + defer func() { + close(done) + if lastErr == nil { + apiresp.GinSuccess(c, nil) + } else { + m.putErr(c, lastErr) + } + }() + doAppend := func() { + if end == false && len(packet) == 0 { + return + } + ctx, cancel := context.WithTimeout(c, time.Second*10) + defer cancel() + req := &msg.AppendStreamMsgReq{ + ConversationID: conversationID, + ClientMsgID: clientMsgID, + StartIndex: int64(index), + Packets: packet, + End: end, + } + _, lastErr = m.Client.AppendStreamMsg(ctx, req) + if lastErr == nil { + log.ZDebug(ctx, "AppendStreamMsg ok", "clientMsgID", clientMsgID) + index += len(packet) + packet = packet[:0] + errCount = 0 + return + } + errCount++ + if errs.ErrRecordNotFound.Is(lastErr) { + log.ZWarn(c, "msg not found", nil, "clientMsgID", clientMsgID) + return + } else if errs.ErrNoPermission.Is(lastErr) { + log.ZError(c, "msg permission error", nil, "clientMsgID", clientMsgID) + return + } else { + log.ZError(c, "append stream msg failed", lastErr, "clientMsgID", clientMsgID, "errCount", errCount) + time.Sleep(time.Millisecond * 50 * time.Duration(errCount)) + } + } + for errCount < 10 { + select { + case s, ok := <-streamCh: + if ok { + packet = append(packet, s) + } + if !ok { + end = true + } + doAppend() + if end == true && lastErr == nil { + return + } + } + } +} + +func NewUTF8Reader(r io.Reader) io.Reader { + return &UTF8Reader{ + r: bufio.NewReaderSize(r, 512), + } +} + +type UTF8Reader struct { + r *bufio.Reader + buf bytes.Buffer +} + +func (r *UTF8Reader) Read(b []byte) (int, error) { + for { + n, err := r.r.Read(b) + if err != nil { + return 0, err + } + r.buf.Write(b[:n]) + data := r.buf.Bytes() + minIndex := min(len(b), len(data)) + if minIndex == 0 { + continue + } + for i := minIndex; i > 0; i-- { + if utf8.Valid(data[:i]) { + n, err := r.buf.Read(b[:i]) + if err != nil { + return 0, err + } + if n != i { + return 0, fmt.Errorf("invalid UTF-8 encoding") + } + return n, nil + } + } + } +} + +func (m *MessageApi) putErr(c *gin.Context, err error) { + c.JSON(http.StatusOK, apiresp.ParseError(err)) +} diff --git a/internal/rpc/msg/modify.go b/internal/rpc/msg/modify.go new file mode 100644 index 000000000..205375dbb --- /dev/null +++ b/internal/rpc/msg/modify.go @@ -0,0 +1,167 @@ +package msg + +import ( + "context" + "encoding/json" + "fmt" + "time" + + "github.com/openimsdk/open-im-server/v3/pkg/common/servererrs" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/model" + "github.com/openimsdk/open-im-server/v3/pkg/msgprocessor" + "github.com/openimsdk/protocol/constant" + msgpb "github.com/openimsdk/protocol/msg" + "github.com/openimsdk/protocol/sdkws" + "github.com/openimsdk/tools/errs" + "github.com/openimsdk/tools/log" + "github.com/openimsdk/tools/mcontext" + "github.com/openimsdk/tools/utils/datautil" +) + +func (m *msgServer) getModifyRawMessage(ctx context.Context, req *msgpb.ModifyMessageReq) (*model.MsgDataModel, error) { + opUserID := mcontext.GetOpUserID(ctx) + msgs, err := m.MsgDatabase.GetMessageBySeqsDB(ctx, req.ConversationID, opUserID, []int64{req.Seq}) + if err != nil { + return nil, err + } + if len(msgs) == 0 { + return nil, errs.ErrRecordNotFound.WrapMsg("msg seq not found") + } + val := msgs[0] + if val == nil || val.Msg == nil || val.Msg.Status == constant.MsgStatusHasDeleted { + return nil, servererrs.ErrRecordNotFound.WrapMsg("msg already delete") + } + if val.Revoke != nil { + return nil, servererrs.ErrMsgAlreadyRevoke.WrapMsg("msg already revoke") + } + msgData := val.Msg + if req.OldContent != "" { + if req.OldContent != msgData.Content { + return nil, servererrs.ErrArgs.WrapMsg("old msg content not match") + } + } + if req.NewContent == msgData.Content { + return nil, errs.ErrArgs.WrapMsg("new content same as old content") + } + if datautil.Contain(opUserID, m.config.Share.IMAdminUser.UserIDs...) { + return msgData, nil + } + isGroup := msgprocessor.IsGroupConversationID(req.ConversationID) + if !isGroup { + if msgData.SendID != opUserID { + return nil, servererrs.ErrNoPermission.WrapMsg("no permission") + } + return msgData, nil + } + groupID := msgData.GroupID + if groupID == "" { + groupID = msgData.RecvID + } + groupInfo, err := m.GroupLocalCache.GetGroupInfo(ctx, groupID) + if err != nil { + return nil, err + } + if groupInfo.Status == constant.GroupStatusDismissed { + return nil, servererrs.ErrDismissedAlready.Wrap() + } + var memberUserIDs []string + if msgData.SendID == opUserID { + memberUserIDs = []string{opUserID} + } else { + memberUserIDs = []string{opUserID, msgData.SendID} + } + members, err := m.GroupLocalCache.GetGroupMemberInfoMap(ctx, groupID, memberUserIDs) + if err != nil { + return nil, err + } + opMember, ok := members[opUserID] + if !ok { + return nil, servererrs.ErrNoPermission.WrapMsg("opUser no in group") + } + if msgData.SendID == opUserID { + return msgData, nil + } + if opMember.RoleLevel <= constant.GroupOrdinaryUsers { + return nil, errs.ErrNoPermission.WrapMsg("no permission update other user msg") + } + var sendRoleLevel int32 + if sendMember, ok := members[msgData.SendID]; ok { + sendRoleLevel = sendMember.RoleLevel + } + if sendRoleLevel >= opMember.RoleLevel { + return nil, errs.ErrNoPermission.WrapMsg("no permission update other user msg") + } + return msgData, nil +} + +func (m *msgServer) ModifyMessage(ctx context.Context, req *msgpb.ModifyMessageReq) (*msgpb.ModifyMessageResp, error) { + lockKey := fmt.Sprintf("MODIFYMESSAGE:%s:%d", req.ConversationID, req.Seq) + lockValue, err := m.lock.Lock(ctx, lockKey, time.Second*30) + if err != nil { + return nil, err + } + defer m.lock.Unlock(ctx, lockKey, lockValue) + msg, err := m.getModifyRawMessage(ctx, req) + if err != nil { + return nil, err + } + var attachedInfo map[string]json.RawMessage + if msg.AttachedInfo != "" && msg.AttachedInfo != "null" && msg.AttachedInfo != "{}" { + if err = json.Unmarshal([]byte(msg.AttachedInfo), &attachedInfo); err != nil { + log.ZWarn(ctx, "json.Unmarshal", err, "attachedInfo", msg.AttachedInfo) + } + } + if attachedInfo == nil { + attachedInfo = make(map[string]json.RawMessage) + } + const modifyAttachedKey = "lastModified" + type LastModified struct { + UserID string `json:"userID"` // last modified user ID + ModifiedTime int64 `json:"modifiedTime"` // last modified time + ModifiedCount int64 `json:"modifiedCount"` // last modified count + } + var modifyValue LastModified + if val := attachedInfo[modifyAttachedKey]; len(val) > 0 { + if err = json.Unmarshal(val, &modifyValue); err != nil { + return nil, errs.WrapMsg(err, "json.Unmarshal modifyValue", "val", val) + } + if modifyValue.ModifiedCount < 1 { + modifyValue.ModifiedCount = 1 + } + } + modifyValue.ModifiedCount++ + modifyValue.ModifiedTime = time.Now().UnixMilli() + modifyValue.UserID = mcontext.GetOpUserID(ctx) + modifyVal, err := json.Marshal(&modifyValue) + if err != nil { + return nil, err + } + attachedInfo[modifyAttachedKey] = modifyVal + attached, err := json.Marshal(attachedInfo) + if err != nil { + return nil, errs.ErrInternalServer.WrapMsg("json.Marshal attachedInfo", "attachedInfo", attachedInfo) + } + msg.Content = req.NewContent + msg.AttachedInfo = string(attached) + if err := m.MsgDatabase.UpdateMsg(ctx, req.ConversationID, msg); err != nil { + return nil, err + } + tips := &sdkws.ModifyMsgTips{ + ConversationID: req.ConversationID, + Seq: req.Seq, + ClientMsgID: msg.ClientMsgID, + NewContent: req.NewContent, + ModifiedTime: modifyValue.ModifiedTime, + ModifiedCount: modifyValue.ModifiedCount, + UserID: modifyValue.UserID, + } + recvID := msg.GroupID + if recvID == "" { + recvID = msg.RecvID + } + m.notificationSender.NotificationWithSessionType(ctx, msg.SendID, recvID, constant.ModifyMessageNotification, msg.SessionType, tips) + return &msgpb.ModifyMessageResp{ + ModifiedTime: modifyValue.ModifiedTime, + ModifiedCount: modifyValue.ModifiedCount, + }, nil +} diff --git a/internal/rpc/msg/notification.go b/internal/rpc/msg/notification.go index 0418823d6..0daafbe6c 100644 --- a/internal/rpc/msg/notification.go +++ b/internal/rpc/msg/notification.go @@ -48,3 +48,7 @@ func (m *MsgNotificationSender) MarkAsReadNotification(ctx context.Context, conv } m.NotificationWithSessionType(ctx, sendID, recvID, constant.HasReadReceipt, sessionType, tips) } + +func (m *MsgNotificationSender) StreamMsgNotification(ctx context.Context, sendID string, recvID string, sessionType int32, tips *sdkws.StreamMsgTips) { + m.NotificationWithSessionType(ctx, sendID, recvID, constant.StreamMsgNotification, sessionType, tips) +} diff --git a/internal/rpc/msg/send.go b/internal/rpc/msg/send.go index 18ad9cc56..60d0fe996 100644 --- a/internal/rpc/msg/send.go +++ b/internal/rpc/msg/send.go @@ -54,6 +54,11 @@ func (m *msgServer) SendMsg(ctx context.Context, req *pbmsg.SendMsgReq) (*pbmsg. func (m *msgServer) sendMsg(ctx context.Context, req *pbmsg.SendMsgReq, before **sdkws.MsgData) (*pbmsg.SendMsgResp, error) { m.encapsulateMsgData(req.MsgData) + if req.MsgData.ContentType == constant.Stream { + if err := m.createStreamMsgHandler(ctx, req.MsgData); err != nil { + return nil, err + } + } switch req.MsgData.SessionType { case constant.SingleChatType: return m.sendMsgSingleChat(ctx, req, before) diff --git a/internal/rpc/msg/server.go b/internal/rpc/msg/server.go index e7c589b8c..cb67ae4bb 100644 --- a/internal/rpc/msg/server.go +++ b/internal/rpc/msg/server.go @@ -24,6 +24,7 @@ import ( "github.com/openimsdk/open-im-server/v3/pkg/rpcli" "github.com/openimsdk/open-im-server/v3/pkg/common/config" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/redis" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/controller" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database/mgo" @@ -69,6 +70,8 @@ type msgServer struct { config *Config // Global configuration settings. webhookClient *webhook.Client conversationClient *rpcli.ConversationClient + lock cache.Lock + StreamMsgDatabase controller.StreamMsgDatabase adminUserIDs []string } @@ -139,6 +142,8 @@ func Start(ctx context.Context, config *Config, client discovery.SvcDiscoveryReg config: config, webhookClient: webhook.NewWebhookClient(config.WebhooksConfig.URL), conversationClient: conversationClient, + lock: redis.NewLock(rdb), + StreamMsgDatabase: controller.NewStreamMsgDatabase(redis.NewStreamMsg(rdb)), adminUserIDs: config.Share.IMAdminUser.UserIDs, } diff --git a/internal/rpc/msg/stream_msg.go b/internal/rpc/msg/stream_msg.go new file mode 100644 index 000000000..ee237e067 --- /dev/null +++ b/internal/rpc/msg/stream_msg.go @@ -0,0 +1,193 @@ +package msg + +import ( + "context" + "encoding/json" + "time" + + "github.com/openimsdk/open-im-server/v3/pkg/apistruct" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" + "github.com/openimsdk/open-im-server/v3/pkg/msgprocessor" + "github.com/openimsdk/protocol/constant" + "github.com/openimsdk/protocol/msg" + "github.com/openimsdk/protocol/sdkws" + "github.com/openimsdk/tools/errs" + "github.com/openimsdk/tools/log" + "github.com/openimsdk/tools/mcontext" +) + +const ( + StreamTimeoutEnd = time.Minute * 10 + StreamTimeoutEndMillisecond = int64(StreamTimeoutEnd / time.Millisecond) +) + +func (m *msgServer) createStreamMsgHandler(ctx context.Context, msgData *sdkws.MsgData) error { + var elem apistruct.StreamMsgElem + if err := json.Unmarshal(msgData.Content, &elem); err != nil { + return errs.ErrArgs.WrapMsg("stream msg content is invalid", "content", string(msgData.Content)) + } + conversationID := msgprocessor.GetConversationIDByMsg(msgData) + if _, err := m.StreamMsgDatabase.GetStreamMsg(ctx, conversationID, msgData.ClientMsgID); err != nil { + if !errs.ErrRecordNotFound.Is(err) { + return err + } + } + streamMsg := &cache.StreamMsg{ + SendUserID: msgData.SendID, + RecvID: msgData.RecvID, + SessionType: msgData.SessionType, + UpdateTime: time.Now().UnixMilli(), + StreamType: elem.Type, + StreamContent: elem.Content, + } + switch msgData.SessionType { + case constant.ReadGroupChatType, constant.WriteGroupChatType: + streamMsg.RecvID = msgData.GroupID + } + return m.StreamMsgDatabase.CreateStreamMsg(ctx, msgprocessor.GetConversationIDByMsg(msgData), msgData.ClientMsgID, streamMsg) +} + +func (m *msgServer) AppendStreamMsg(ctx context.Context, req *msg.AppendStreamMsgReq) (*msg.AppendStreamMsgResp, error) { + end, err := m.StreamMsgDatabase.GetStreamMsgEnd(ctx, req.ConversationID, req.ClientMsgID) + if err != nil { + return nil, err + } + if end { + return nil, errs.ErrNoPermission.WrapMsg("stream msg is end") + } + res, err := m.StreamMsgDatabase.AppendStreamMsg(ctx, req.ConversationID, req.ClientMsgID, int(req.StartIndex), req.Packets, req.End, req.End) + if err != nil { + return nil, err + } + tips := &sdkws.StreamMsgTips{ + ConversationID: req.ConversationID, + ClientMsgID: req.ClientMsgID, + StartIndex: req.StartIndex, + Packets: req.Packets, + End: req.End, + } + m.msgNotificationSender.StreamMsgNotification(ctx, res.SendUserID, res.RecvID, res.SessionType, tips) + if req.End { + m.modifyStreamMessage(ctx, req.ConversationID, req.ClientMsgID, res) + } + return &msg.AppendStreamMsgResp{}, nil +} + +func (m *msgServer) modifyStreamMessage(ctx context.Context, conversationID string, clientMsgID string, res *cache.StreamMsg) { + packets := make([]string, 0, len(res.Packets)) + for i := int64(0); ; i++ { + data, ok := res.Packets[i] + if !ok { + break + } + packets = append(packets, data) + } + content, err := json.Marshal(&apistruct.StreamMsgElem{ + Type: res.StreamType, + Content: res.StreamContent, + Packets: packets, + End: res.End, + Deadline: res.UpdateTime, + }) + if err != nil { + log.ZError(ctx, "modifyStreamMessage json.Marshal", err, "conversationID", conversationID, "clientMsgID", clientMsgID) + return + } + req := &msg.ModifyMessageReq{ + ConversationID: conversationID, + NewContent: string(content), + } + modifyMessage := func() error { + ctx, cancel := context.WithTimeout(ctx, time.Second*10) + defer cancel() + if req.Seq == 0 { + req.Seq, err = m.MsgDatabase.GetMessageSeq(ctx, conversationID, clientMsgID) + if err != nil { + return err + } + } + if _, err := m.ModifyMessage(ctx, req); err != nil { + return err + } + return nil + } + + if err := modifyMessage(); err != nil { + log.ZError(ctx, "sync modifyStreamMessage", err, "conversationID", conversationID, "content", string(content)) + ctx = context.WithoutCancel(ctx) + go func() { + for i := 1; i <= 10; i++ { + if err := modifyMessage(); err == nil { + log.ZDebug(ctx, "async modifyStreamMessage success", "conversationID", conversationID, "content", string(content), "count", i) + return + } else { + log.ZError(ctx, "modifyStreamMessage", err, "conversationID", conversationID, "content", string(content), "count", i) + time.Sleep(time.Second * time.Duration(i)) + } + } + }() + } +} + +func (m *msgServer) GetStreamMsg(ctx context.Context, req *msg.GetStreamMsgReq) (*msg.GetStreamMsgResp, error) { + value, err := m.StreamMsgDatabase.GetStreamMsg(ctx, req.ConversationID, req.ClientMsgID) + if err == nil { + resp := msg.GetStreamMsgResp{ + UserID: value.SendUserID, + Packets: make([]string, 0, len(value.Packets)), + End: value.End, + } + for i := int64(0); ; i++ { + data, ok := value.Packets[i] + if !ok { + break + } + resp.Packets = append(resp.Packets, data) + } + if resp.End { + resp.DeadlineTime = value.UpdateTime + } else { + if now := time.Now().UnixMilli(); now-value.UpdateTime >= StreamTimeoutEndMillisecond { + resp.DeadlineTime = now + StreamTimeoutEndMillisecond + resp.End = true + } + } + return &resp, nil + } else if !errs.ErrRecordNotFound.Is(err) { + return nil, err + } + if req.Seq <= 0 || errs.ErrRecordNotFound.Is(err) == false { + return nil, err + } + msgs, err := m.MsgDatabase.GetMessageBySeqs(ctx, req.ConversationID, mcontext.GetOpUserID(ctx), []int64{req.Seq}) + if err != nil { + return nil, err + } + if len(msgs) == 0 || msgs[0] == nil { + return nil, errs.ErrRecordNotFound.WrapMsg("stream message not found") + } + msgData := msgs[0] + if msgData.ClientMsgID != req.ClientMsgID { + return nil, errs.ErrRecordNotFound.WrapMsg("stream message id not match") + } + if msgData.ContentType != constant.Stream { + return nil, errs.ErrNoPermission.WrapMsg("stream message content type not match") + } + var elem apistruct.StreamMsgElem + if len(msgData.Content) > 0 { + if err := json.Unmarshal(msgData.Content, &elem); err != nil { + log.ZError(ctx, "stream msg unmarshal", err, "content", string(msgData.Content), "conversationID", req.ConversationID, "seq", req.Seq) + } + } + resp := &msg.GetStreamMsgResp{ + UserID: msgData.SendID, + Packets: elem.Packets, + End: elem.End, + DeadlineTime: elem.Deadline, + } + if !resp.End { + resp.End = true + resp.DeadlineTime = msgData.SendTime + StreamTimeoutEndMillisecond + } + return resp, nil +} diff --git a/pkg/apistruct/msg.go b/pkg/apistruct/msg.go index 0e1b356a1..2333ca2dd 100644 --- a/pkg/apistruct/msg.go +++ b/pkg/apistruct/msg.go @@ -90,8 +90,11 @@ type MarkdownTextElem struct { } type StreamMsgElem struct { - Type string `mapstructure:"type" validate:"required"` - Content string `mapstructure:"content" validate:"required"` + Type string `mapstructure:"type" json:"type"` + Content string `mapstructure:"content" json:"content"` + Packets []string `mapstructure:"packets" json:"packets"` + End bool `mapstructure:"end" json:"end"` + Deadline int64 `mapstructure:"deadline" json:"deadline"` } type RevokeElem struct { diff --git a/pkg/common/storage/cache/cachekey/msg.go b/pkg/common/storage/cache/cachekey/msg.go index ac449df38..6f0eb90d9 100644 --- a/pkg/common/storage/cache/cachekey/msg.go +++ b/pkg/common/storage/cache/cachekey/msg.go @@ -21,6 +21,7 @@ import ( const ( sendMsgFailedFlag = "SEND_MSG_FAILED_FLAG:" messageCache = "MSG_CACHE:" + messageSeq = "MSG_SEQ:" ) func GetMsgCacheKey(conversationID string, seq int64) string { @@ -30,3 +31,7 @@ func GetMsgCacheKey(conversationID string, seq int64) string { func GetSendMsgKey(id string) string { return sendMsgFailedFlag + id } + +func GetMsgSeqKey(conversationID string, clientMsgID string) string { + return messageSeq + conversationID + ":" + clientMsgID +} diff --git a/pkg/common/storage/cache/cachekey/stream_msg.go b/pkg/common/storage/cache/cachekey/stream_msg.go new file mode 100644 index 000000000..613296676 --- /dev/null +++ b/pkg/common/storage/cache/cachekey/stream_msg.go @@ -0,0 +1,7 @@ +package cachekey + +const streamMessageCache = "STREAM_MSG:" + +func GetStreamMsgKey(conversationID string, clientMsgID string) string { + return streamMessageCache + conversationID + ":" + clientMsgID +} diff --git a/pkg/common/storage/cache/lock.go b/pkg/common/storage/cache/lock.go new file mode 100644 index 000000000..b867dc70a --- /dev/null +++ b/pkg/common/storage/cache/lock.go @@ -0,0 +1,11 @@ +package cache + +import ( + "context" + "time" +) + +type Lock interface { + Lock(ctx context.Context, key string, timeout time.Duration) (string, error) + Unlock(ctx context.Context, key, value string) +} diff --git a/pkg/common/storage/cache/msg.go b/pkg/common/storage/cache/msg.go index 271ed19fe..09e688956 100644 --- a/pkg/common/storage/cache/msg.go +++ b/pkg/common/storage/cache/msg.go @@ -16,6 +16,7 @@ package cache import ( "context" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/model" ) @@ -26,4 +27,5 @@ type MsgCache interface { GetMessageBySeqs(ctx context.Context, conversationID string, seqs []int64) ([]*model.MsgInfoModel, error) DelMessageBySeqs(ctx context.Context, conversationID string, seqs []int64) error SetMessageBySeqs(ctx context.Context, conversationID string, msgs []*model.MsgInfoModel) error + GetMessageSeq(ctx context.Context, conversationID string, clientMsgID string) (int64, error) } diff --git a/pkg/common/storage/cache/redis/lock.go b/pkg/common/storage/cache/redis/lock.go new file mode 100644 index 000000000..ad4a887d5 --- /dev/null +++ b/pkg/common/storage/cache/redis/lock.go @@ -0,0 +1,60 @@ +package redis + +import ( + "context" + "encoding/hex" + "time" + + "github.com/google/uuid" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" + "github.com/openimsdk/tools/errs" + "github.com/openimsdk/tools/log" + "github.com/redis/go-redis/v9" +) + +const lockPrefix = "LOCK:" + +func NewLock(rdb redis.UniversalClient) cache.Lock { + return &redisLock{rdb: rdb} +} + +type redisLock struct { + rdb redis.UniversalClient +} + +func (x *redisLock) Lock(ctx context.Context, key string, timeout time.Duration) (string, error) { + uid, err := uuid.NewUUID() + if err != nil { + return "", err + } + if timeout < time.Second { + timeout = time.Minute * 2 + } + value := hex.EncodeToString(uid[:]) + key = lockPrefix + key + for { + ok, err := x.rdb.SetNX(ctx, key, value, timeout).Result() + if err != nil { + return "", errs.WrapMsg(err, "get redis lock", "key", key) + } + if ok { + return value, nil + } + timer := time.NewTimer(50 * time.Millisecond) + select { + case <-ctx.Done(): + timer.Stop() + return "", context.Cause(ctx) + case <-timer.C: + } + } +} + +func (x *redisLock) Unlock(ctx context.Context, key, value string) { + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second) + defer cancel() + script := "\nlocal value = redis.call(\"GET\", KEYS[1])\nif value == ARGV[1] then\n return redis.call(\"DEL\", KEYS[1])\nend\nreturn 0" + if err := x.rdb.Eval(ctx, script, []string{lockPrefix + key}, value).Err(); err != nil { + log.ZWarn(ctx, "unlock redis lock", err, "key", key) + } +} diff --git a/pkg/common/storage/cache/redis/msg.go b/pkg/common/storage/cache/redis/msg.go index dfe6ca04d..64e982099 100644 --- a/pkg/common/storage/cache/redis/msg.go +++ b/pkg/common/storage/cache/redis/msg.go @@ -9,6 +9,7 @@ import ( "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/database" "github.com/openimsdk/open-im-server/v3/pkg/common/storage/model" + "github.com/openimsdk/protocol/constant" "github.com/openimsdk/tools/errs" "github.com/openimsdk/tools/utils/datautil" "github.com/redis/go-redis/v9" @@ -89,6 +90,15 @@ func (c *msgCache) SetMessageBySeqs(ctx context.Context, conversationID string, if err := c.rcClient.GetClient().RawSet(ctx, cachekey.GetMsgCacheKey(conversationID, msg.Msg.Seq), string(data), msgCacheTimeout); err != nil { return err } + if msg.Msg.ContentType == constant.Stream { + if err := c.rcClient.GetRedis().Set(ctx, cachekey.GetMsgSeqKey(conversationID, msg.Msg.ClientMsgID), msg.Msg.Seq, msgCacheTimeout).Err(); err != nil { + return err + } + } } return nil } + +func (c *msgCache) GetMessageSeq(ctx context.Context, conversationID string, clientMsgID string) (int64, error) { + return c.rcClient.GetRedis().Get(ctx, cachekey.GetMsgSeqKey(conversationID, clientMsgID)).Int64() +} diff --git a/pkg/common/storage/cache/redis/stream_msg.go b/pkg/common/storage/cache/redis/stream_msg.go new file mode 100644 index 000000000..7f4062360 --- /dev/null +++ b/pkg/common/storage/cache/redis/stream_msg.go @@ -0,0 +1,140 @@ +package redis + +import ( + "context" + "fmt" + "strconv" + "strings" + "time" + + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" + "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache/cachekey" + "github.com/openimsdk/tools/errs" + "github.com/redis/go-redis/v9" +) + +func NewStreamMsg(rdb redis.UniversalClient) cache.StreamMsgCache { + return &streamMsg{rdb: rdb} +} + +type streamMsg struct { + rdb redis.UniversalClient +} + +func (x *streamMsg) getMsgKey(conversationID string, clientMsgID string) string { + return cachekey.GetStreamMsgKey(conversationID, clientMsgID) +} + +func (x *streamMsg) CreateStreamMsg(ctx context.Context, conversationID string, clientMsgID string, msg *cache.StreamMsg) error { + key := x.getMsgKey(conversationID, clientMsgID) + pipeline := x.rdb.Pipeline() + pipeline.HSet(ctx, key, "sendUserID", msg.SendUserID) + pipeline.HSet(ctx, key, "recvID", msg.RecvID) + pipeline.HSet(ctx, key, "sessionType", strconv.Itoa(int(msg.SessionType))) + pipeline.HSet(ctx, key, "updateTime", time.Now().UnixMilli()) + pipeline.HSet(ctx, key, "isEnd", false) + pipeline.HSet(ctx, key, "streamType", msg.StreamType) + pipeline.HSet(ctx, key, "streamContent", msg.StreamContent) + pipeline.Expire(ctx, key, 24*time.Hour) + _, err := pipeline.Exec(ctx) + return err +} + +func (x *streamMsg) AppendStreamMsg(ctx context.Context, conversationID string, clientMsgID string, startIndex int, packets []string, end bool, retPacket bool) (*cache.StreamMsg, error) { + key := x.getMsgKey(conversationID, clientMsgID) + var mapCmd *redis.MapStringStringCmd + var sliceCmd *redis.SliceCmd + pipeline := x.rdb.Pipeline() + for i, packet := range packets { + pipeline.HSet(ctx, key, "i_"+strconv.Itoa(startIndex+i), packet) + } + pipeline.HSet(ctx, key, "isEnd", end) + pipeline.HSet(ctx, key, "updateTime", time.Now().UnixMilli()) + pipeline.Expire(ctx, key, 24*time.Hour) + if retPacket { + mapCmd = pipeline.HGetAll(ctx, key) + } else { + sliceCmd = pipeline.HMGet(ctx, key, "sendUserID", "recvID", "sessionType") + } + if _, err := pipeline.Exec(ctx); err != nil { + return nil, err + } + var data map[string]string + var err error + if retPacket { + data, err = mapCmd.Result() + } else { + arr, resultErr := sliceCmd.Result() + if resultErr != nil { + return nil, resultErr + } + if len(arr) != 3 || arr[0] == nil || arr[1] == nil || arr[2] == nil { + return nil, errs.ErrRecordNotFound.WrapMsg("stream message not found") + } + data = map[string]string{ + "sendUserID": fmt.Sprint(arr[0]), "recvID": fmt.Sprint(arr[1]), "sessionType": fmt.Sprint(arr[2]), + } + } + if err != nil { + return nil, err + } + return x.mapToStreamMsg(data, retPacket) +} + +func (x *streamMsg) GetStreamMsg(ctx context.Context, conversationID string, clientMsgID string) (*cache.StreamMsg, error) { + data, err := x.rdb.HGetAll(ctx, x.getMsgKey(conversationID, clientMsgID)).Result() + if err != nil { + return nil, err + } + if len(data) == 0 { + return nil, errs.ErrRecordNotFound.WrapMsg("stream message not found") + } + return x.mapToStreamMsg(data, true) +} + +func (x *streamMsg) mapToStreamMsg(data map[string]string, full bool) (*cache.StreamMsg, error) { + sessionType, err := strconv.ParseInt(data["sessionType"], 10, 32) + if err != nil { + return nil, err + } + if !full { + return &cache.StreamMsg{SendUserID: data["sendUserID"], RecvID: data["recvID"], SessionType: int32(sessionType)}, nil + } + end, err := strconv.ParseBool(data["isEnd"]) + if err != nil { + return nil, err + } + updateTime, err := strconv.ParseInt(data["updateTime"], 10, 64) + if err != nil { + return nil, err + } + msg := cache.StreamMsg{ + SendUserID: data["sendUserID"], RecvID: data["recvID"], SessionType: int32(sessionType), + StreamType: data["streamType"], StreamContent: data["streamContent"], UpdateTime: updateTime, + Packets: make(map[int64]string), End: end, + } + var maxIndex int64 = -1 + for indexStr, value := range data { + if !strings.HasPrefix(indexStr, "i_") { + continue + } + index, err := strconv.ParseInt(strings.TrimPrefix(indexStr, "i_"), 10, 64) + if err != nil || index < 0 { + return nil, errs.ErrInternalServer.WrapMsg("packet index is invalid", "index", indexStr) + } + msg.Packets[index] = value + if maxIndex < index { + maxIndex = index + } + } + for i := int64(0); i <= maxIndex; i++ { + if _, ok := msg.Packets[i]; !ok { + return nil, errs.ErrInternalServer.WrapMsg("packet index is not continuous", "index", i) + } + } + return &msg, nil +} + +func (x *streamMsg) GetStreamMsgEnd(ctx context.Context, conversationID string, clientMsgID string) (bool, error) { + return x.rdb.HGet(ctx, x.getMsgKey(conversationID, clientMsgID), "isEnd").Bool() +} diff --git a/pkg/common/storage/cache/stream_msg.go b/pkg/common/storage/cache/stream_msg.go new file mode 100644 index 000000000..eec55df6b --- /dev/null +++ b/pkg/common/storage/cache/stream_msg.go @@ -0,0 +1,21 @@ +package cache + +import "context" + +type StreamMsg struct { + SendUserID string + RecvID string + SessionType int32 + StreamType string + StreamContent string + Packets map[int64]string + End bool + UpdateTime int64 +} + +type StreamMsgCache interface { + CreateStreamMsg(ctx context.Context, conversationID string, clientMsgID string, msg *StreamMsg) error + AppendStreamMsg(ctx context.Context, conversationID string, clientMsgID string, startIndex int, packets []string, end bool, retPacket bool) (*StreamMsg, error) + GetStreamMsg(ctx context.Context, conversationID string, clientMsgID string) (*StreamMsg, error) + GetStreamMsgEnd(ctx context.Context, conversationID string, clientMsgID string) (bool, error) +} diff --git a/pkg/common/storage/controller/msg.go b/pkg/common/storage/controller/msg.go index f833008e8..26742fefa 100644 --- a/pkg/common/storage/controller/msg.go +++ b/pkg/common/storage/controller/msg.go @@ -101,6 +101,11 @@ type CommonMsgDatabase interface { GetLastMessageSeqByTime(ctx context.Context, conversationID string, time int64) (int64, error) GetLastMessage(ctx context.Context, conversationIDS []string, userID string) (map[string]*sdkws.MsgData, error) + + GetMessageBySeqs(ctx context.Context, conversationID string, userID string, seqs []int64) ([]*sdkws.MsgData, error) + GetMessageBySeqsDB(ctx context.Context, conversationID string, userID string, seqs []int64) ([]*model.MsgInfoModel, error) + GetMessageSeq(ctx context.Context, conversationID string, clientMsgID string) (int64, error) + UpdateMsg(ctx context.Context, conversationID string, msg *model.MsgDataModel) error } func NewCommonMsgDatabase(msgDocModel database.Msg, msg cache.MsgCache, seqUser cache.SeqUser, seqConversation cache.SeqConversationCache, producer mq.Producer) CommonMsgDatabase { @@ -822,6 +827,29 @@ func (db *commonMsgDatabase) GetMessageBySeqs(ctx context.Context, conversationI return res, nil } +func (db *commonMsgDatabase) GetMessageBySeqsDB(ctx context.Context, conversationID string, userID string, seqs []int64) ([]*model.MsgInfoModel, error) { + msgs, err := db.msgCache.GetMessageBySeqs(ctx, conversationID, seqs) + if err != nil { + return nil, err + } + db.handlerDeleteAndRevoked(ctx, userID, msgs) + db.handlerQuote(ctx, userID, conversationID, msgs) + return msgs, nil +} + +func (db *commonMsgDatabase) GetMessageSeq(ctx context.Context, conversationID string, clientMsgID string) (int64, error) { + return db.msgCache.GetMessageSeq(ctx, conversationID, clientMsgID) +} + +func (db *commonMsgDatabase) UpdateMsg(ctx context.Context, conversationID string, msg *model.MsgDataModel) error { + docID := db.msgTable.GetDocID(conversationID, msg.Seq) + index := db.msgTable.GetMsgIndex(msg.Seq) + if _, err := db.msgDocDatabase.UpdateMsg(ctx, docID, index, "msg", msg); err != nil { + return err + } + return db.msgCache.DelMessageBySeqs(ctx, conversationID, []int64{msg.Seq}) +} + func (db *commonMsgDatabase) GetLastMessage(ctx context.Context, conversationIDs []string, userID string) (map[string]*sdkws.MsgData, error) { res := make(map[string]*sdkws.MsgData) for _, conversationID := range conversationIDs { diff --git a/pkg/common/storage/controller/stream_msg.go b/pkg/common/storage/controller/stream_msg.go new file mode 100644 index 000000000..96c1c2af5 --- /dev/null +++ b/pkg/common/storage/controller/stream_msg.go @@ -0,0 +1,15 @@ +package controller + +import "github.com/openimsdk/open-im-server/v3/pkg/common/storage/cache" + +type StreamMsgDatabase interface { + cache.StreamMsgCache +} + +func NewStreamMsgDatabase(db cache.StreamMsgCache) StreamMsgDatabase { + return &streamMsgDatabase{db} +} + +type streamMsgDatabase struct { + cache.StreamMsgCache +} diff --git a/pkg/notification/msg.go b/pkg/notification/msg.go index ba8a9185a..797b68b29 100644 --- a/pkg/notification/msg.go +++ b/pkg/notification/msg.go @@ -74,9 +74,10 @@ func newContentTypeConf(conf *config.Notification) map[int32]config.Notification constant.ConversationUnreadNotification: conf.ConversationChanged, constant.ConversationPrivateChatNotification: conf.ConversationSetPrivate, // msg - constant.MsgRevokeNotification: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, - constant.HasReadReceipt: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, - constant.DeleteMsgsNotification: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, + constant.MsgRevokeNotification: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, + constant.HasReadReceipt: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, + constant.DeleteMsgsNotification: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, + constant.ModifyMessageNotification: {IsSendMsg: false, ReliabilityLevel: constant.ReliableNotificationNoMsg}, } } From 7fb7d55166f244d673c30fd5609aac7d1ccb6c3a Mon Sep 17 00:00:00 2001 From: withchao <993506633@qq.com> Date: Tue, 4 Aug 2026 14:15:43 +0800 Subject: [PATCH 19/19] chore: update openimsdk/tools dependency to v0.0.50-alpha.122 --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index 61349964e..e3f52757a 100644 --- a/go.mod +++ b/go.mod @@ -13,7 +13,7 @@ require ( github.com/grpc-ecosystem/go-grpc-prometheus v1.2.0 github.com/mitchellh/mapstructure v1.5.0 github.com/openimsdk/protocol v0.0.73-alpha.20 - github.com/openimsdk/tools v0.0.50-alpha.121 + github.com/openimsdk/tools v0.0.50-alpha.122 github.com/pkg/errors v0.9.1 // indirect github.com/prometheus/client_golang v1.18.0 github.com/stretchr/testify v1.11.1 diff --git a/go.sum b/go.sum index f4b070a3f..8ae6bab5b 100644 --- a/go.sum +++ b/go.sum @@ -363,8 +363,8 @@ github.com/openimsdk/gomake v0.0.17 h1:q8haP48VOH45WhJRiLj1YSBJyUFJqD8CTedH65i1Y github.com/openimsdk/gomake v0.0.17/go.mod h1:nnjS8yCtrPJAt1knMbyPiUwCH2gpyBzj/EZAONfUOXg= github.com/openimsdk/protocol v0.0.73-alpha.20 h1:9MnACSi6IKv2iqlxHYUJG9mgt9gyPRHxE7Lq8dxAoxI= github.com/openimsdk/protocol v0.0.73-alpha.20/go.mod h1:WF7EuE55vQvpyUAzDXcqg+B+446xQyEba0X35lTINmw= -github.com/openimsdk/tools v0.0.50-alpha.121 h1:TXKKgtkeMeqIs0vpolbW8rIEngE9xlESq+0NV+FoLH0= -github.com/openimsdk/tools v0.0.50-alpha.121/go.mod h1:I0WESSa7ghPIo9BL+ETlH/qEIbO6+KZioM1jwNuDwz0= +github.com/openimsdk/tools v0.0.50-alpha.122 h1:bKP6hrJ6kGyGUqTGuyqUqPdZxY7mZSsh86aRQJ/6Jt8= +github.com/openimsdk/tools v0.0.50-alpha.122/go.mod h1:I0WESSa7ghPIo9BL+ETlH/qEIbO6+KZioM1jwNuDwz0= github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6Wq+LM= github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs= github.com/pierrec/lz4/v4 v4.1.21 h1:yOVMLb6qSIDP67pl/5F7RepeKYu/VmTyEXvuMI5d9mQ=