From 00f70b14d03bf4719f29ef1caae2bec30f305383 Mon Sep 17 00:00:00 2001 From: tangmingyou Date: Wed, 10 Jan 2024 17:56:44 +0800 Subject: [PATCH] deliver msg --- api/postal.proto | 22 +- cmd/auth/config.toml | 1 + cmd/auth/main.go | 14 +- cmd/chat/config.toml | 33 +++ cmd/chat/main.go | 46 +++++ cmd/gateway_http/config.toml | 24 +++ cmd/gateway_http/main.go | 113 ++++++++++ cmd/gateway_ws/config.toml | 3 + cmd/gateway_ws/main.go | 66 +++++- cmd/mahjong/main.go | 30 +++ go.mod | 12 +- go.sum | 24 ++- internal/auth/data/auth_user_dao.go | 66 ++++++ internal/auth/data/t_user.go | 18 ++ internal/auth/logic/auth_server.go | 109 ++++++++-- internal/chat/logic/chat_server.go | 69 ++++++- internal/gateway_http/config/auth_filter.go | 55 +++++ internal/gateway_http/logic/http_server.go | 1 + internal/gateway_ws/server/ws_server.go | 94 ++++++++- internal/gateway_ws/session/net_account.go | 15 -- internal/gateway_ws/session/net_client.go | 27 ++- internal/gateway_ws/session/net_subject.go | 15 ++ .../postal/logic/postal_cluster_server.go | 77 +++++++ internal/postal/logic/postal_server.go | 194 ++++++++++++++++-- pkg/config/grpc_options.go | 1 - pkg/deliver/deliver.go | 1 - pkg/grpc/client/direct_client_factory.go | 41 ++++ pkg/grpc/discovery/register.go | 27 ++- pkg/grpc/generic/generic_client.go | 34 ++- pkg/grpc/generic/generic_client_factory.go | 11 +- pkg/plugins/cache/cache.go | 46 +++++ pkg/plugins/cache/multi_cache.go | 152 ++++++++++++++ pkg/plugins/cache/redis.go | 52 +++++ pkg/plugins/mq/mq.go | 30 +++ pkg/plugins/mq/nats.go | 141 +++++++++++++ pkg/plugins/mq/nats_jet_stream.go | 190 +++++++++++++++++ pkg/protocol/authorize/authorize.go | 36 ++++ pkg/protocol/deliver/deliver.go | 127 ++++++++++++ pkg/protocol/deliver/options.go | 19 ++ pkg/protocol/session/subject.go | 49 +++++ pkg/utils/cache/cache.go | 37 ++++ pkg/utils/cache/kvcache.go | 59 ++++++ pkg/utils/conver/unit_conver.go | 13 ++ pkg/utils/security/aes.go | 116 +++++++++++ pkg/utils/shutdown/signal.go | 5 +- pkg/utils/strs/random.go | 85 ++++++++ pkg/utils/strs/variable.go | 74 +++++++ 47 files changed, 2357 insertions(+), 117 deletions(-) create mode 100644 cmd/chat/config.toml create mode 100644 cmd/chat/main.go create mode 100644 cmd/gateway_http/config.toml create mode 100644 cmd/gateway_http/main.go create mode 100644 cmd/mahjong/main.go create mode 100644 internal/auth/data/auth_user_dao.go create mode 100644 internal/auth/data/t_user.go create mode 100644 internal/gateway_http/config/auth_filter.go create mode 100644 internal/gateway_http/logic/http_server.go delete mode 100644 internal/gateway_ws/session/net_account.go create mode 100644 internal/gateway_ws/session/net_subject.go create mode 100644 internal/postal/logic/postal_cluster_server.go delete mode 100644 pkg/deliver/deliver.go create mode 100644 pkg/grpc/client/direct_client_factory.go create mode 100644 pkg/plugins/cache/cache.go create mode 100644 pkg/plugins/cache/multi_cache.go create mode 100644 pkg/plugins/cache/redis.go create mode 100644 pkg/plugins/mq/mq.go create mode 100644 pkg/plugins/mq/nats.go create mode 100644 pkg/plugins/mq/nats_jet_stream.go create mode 100644 pkg/protocol/authorize/authorize.go create mode 100644 pkg/protocol/deliver/deliver.go create mode 100644 pkg/protocol/deliver/options.go create mode 100644 pkg/protocol/session/subject.go create mode 100644 pkg/utils/cache/cache.go create mode 100644 pkg/utils/cache/kvcache.go create mode 100644 pkg/utils/security/aes.go create mode 100644 pkg/utils/strs/random.go create mode 100644 pkg/utils/strs/variable.go diff --git a/api/postal.proto b/api/postal.proto index 5071aae..79232e5 100644 --- a/api/postal.proto +++ b/api/postal.proto @@ -3,6 +3,11 @@ import "google/protobuf/empty.proto"; option go_package = "./postal"; +enum DeliverResult { + Success = 0; + ReceiverOffline = 1; +} + service Postal { rpc Deliver(ReqDeliver) returns (ResDeliver); rpc DeliverBatch(ReqDeliverBatch) returns(ResDeliver); @@ -25,7 +30,7 @@ message Message { message ReqDeliver { string receiver = 1; - Message message = 2; + Message msg = 2; // bool sync = 7; // 是否同步阻塞等待投递结果, 默认false立即返回放到队列消费投递 } @@ -70,3 +75,18 @@ message ReqGroupLeave { message ReqGroupDissolve { string gid = 1; } + +// 集群之间接口调用,重定向消息... +service PostalCluster { + // socket不在当前节点,重新投递消息 + rpc Redirect(ReqRedirect) returns(ResDeliver); +} + +message ReqRedirect { + int32 ttl = 1; // 投递一次减一, 到0丢弃 + int32 redirectMethod = 2; // 10 deliver, 11 deliverBatch, 12 deliverGroup + // bool sync = 7; // 是否同步阻塞等待投递结果,否则放到channel顺序投递 + ReqDeliver deliver = 10; + ReqDeliverBatch deliverBatch = 11; + ReqDeliverGroup deliverGroup = 12; +} diff --git a/cmd/auth/config.toml b/cmd/auth/config.toml index 9f48a60..7530b9f 100644 --- a/cmd/auth/config.toml +++ b/cmd/auth/config.toml @@ -1,4 +1,5 @@ [app] +aesTokenKey = "9Nz3Y6DES3msAFndz4QsJAECUIFkf+KKaRa+jRnNALk=" [grpc] address = ":7020" diff --git a/cmd/auth/main.go b/cmd/auth/main.go index 9303dba..8761e3f 100644 --- a/cmd/auth/main.go +++ b/cmd/auth/main.go @@ -3,6 +3,7 @@ package main import ( clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/grpc/grpclog" + "sonet/internal/auth/data" "sonet/internal/auth/logic" "sonet/pkg/config" "sonet/pkg/grpc/discovery" @@ -10,10 +11,14 @@ import ( "sonet/pkg/utils/shutdown" ) +type AuthConfig struct { + AesTokenKey string +} + func main() { grpclog.SetLoggerV2(logger.Logger) - - conf := config.LoadConfig(nil, "cmd/auth") + appConf := &AuthConfig{} + conf := config.LoadConfig(appConf, "cmd/auth") // registry and run... etcdClient, err := clientv3.New(conf.Etcd) @@ -23,7 +28,10 @@ func main() { registry := discovery.NewRegister(etcdClient) shutdown.AddShutdownHook(registry.Stop) - authServer := logic.NewAuthServer() + + // user dao + userDao := data.NewUserDao(config.NewGorm(conf.Gorm)) + authServer := logic.NewAuthServer(appConf.AesTokenKey, userDao) go func() { err = authServer.Run(conf.Grpc, registry) if err != nil { diff --git a/cmd/chat/config.toml b/cmd/chat/config.toml new file mode 100644 index 0000000..10e640b --- /dev/null +++ b/cmd/chat/config.toml @@ -0,0 +1,33 @@ +[app] + +[grpc] +address = ":7030" +maxSendMsgSize = "8Mi" +maxRecvMsgSize = "8Mi" +readBufferSize = "8Ki" +writeBufferSize = "8Ki" + +[grpc.register.attrs] +weight = 100 + +[etcd] +endpoints = ["124.222.131.236:3279"] +username = "root" +password = "sopod@etcd" + +[redis] +Addr = "124.222.131.236:3379" +Password = "sopod@redis#" +DB = 1 +MinIdleConns = 3 + +[gorm] +logMode=true + +[gorm.mysql] +# https://gorm.io/zh_CN/docs/connecting_to_the_database.html +DSN = "root:sopod_mysql2347-@tcp(124.222.131.236:3666)/groups?charset=utf8&parseTime=True&loc=Local" + +[prometheus] +enable = false +port = 7039 diff --git a/cmd/chat/main.go b/cmd/chat/main.go new file mode 100644 index 0000000..e09ce94 --- /dev/null +++ b/cmd/chat/main.go @@ -0,0 +1,46 @@ +package main + +import ( + "context" + clientv3 "go.etcd.io/etcd/client/v3" + "google.golang.org/grpc/grpclog" + "google.golang.org/grpc/resolver" + "sonet/api/gen/chat" + "sonet/internal/chat/logic" + "sonet/pkg/config" + "sonet/pkg/grpc/discovery" + "sonet/pkg/protocol/deliver" + "sonet/pkg/utils/logger" + "sonet/pkg/utils/shutdown" +) + +func main() { + grpclog.SetLoggerV2(logger.Logger) + conf := config.LoadConfig(nil, "cmd/chat") + + // registry and run... + etcdClient, err := clientv3.New(conf.Etcd) + if err != nil { + panic(err) + } + + registry := discovery.NewRegister(etcdClient) + shutdown.AddShutdownHook(registry.Stop) + + etcdResolver := discovery.NewResolver(etcdClient) + resolver.Register(etcdResolver) + + deli := deliver.NewDeliver(chat.Chat_ServiceDesc.ServiceName) + if err := deli.InitWithResolver(context.Background(), etcdResolver); err != nil { + panic(err) + } + chatServer := logic.NewChatServer(deli) + go func() { + err = chatServer.Run(conf.Grpc, registry) + if err != nil { + panic(err) + } + }() + + shutdown.Await() +} diff --git a/cmd/gateway_http/config.toml b/cmd/gateway_http/config.toml new file mode 100644 index 0000000..f644a97 --- /dev/null +++ b/cmd/gateway_http/config.toml @@ -0,0 +1,24 @@ +[app] +port = 7000 +ignoreUrls = [ + "/api/svc/auth/login", + "/api/svc/auth/verify", +] + +[etcd] +endpoints = ["124.222.131.236:3279"] +username = "root" +password = "sopod@etcd" + +[redis] +Addr = "124.222.131.236:3379" +Password = "sopod@redis#" +DB = 1 +MinIdleConns = 3 + +[gorm] +logMode=true + +[prometheus] +enable = false +port = 7039 diff --git a/cmd/gateway_http/main.go b/cmd/gateway_http/main.go new file mode 100644 index 0000000..99455e7 --- /dev/null +++ b/cmd/gateway_http/main.go @@ -0,0 +1,113 @@ +package main + +import ( + "context" + "fmt" + "github.com/gin-gonic/gin" + clientv3 "go.etcd.io/etcd/client/v3" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/grpclog" + "google.golang.org/grpc/resolver" + "net/http" + config2 "sonet/internal/gateway_http/config" + "sonet/pkg/config" + "sonet/pkg/grpc/discovery" + "sonet/pkg/grpc/generic" + "sonet/pkg/utils/logger" + "sonet/pkg/utils/resp" + "sonet/pkg/utils/shutdown" + "sonet/pkg/utils/strs" +) + +type GatewayHttpConfig struct { + Port int + IgnoreUrls []string + AesTokenKey string +} + +func main() { + grpclog.SetLoggerV2(logger.Logger) + + appConf := &GatewayHttpConfig{} + conf := config.LoadConfig(appConf, "cmd/gateway_http") + + // grpc generic client factory + etcdClient, err := clientv3.New(conf.Etcd) + if err != nil { + panic(err) + } + etcdResolver := discovery.NewResolver(etcdClient) + shutdown.AddShutdownHook(etcdResolver.Close) + resolver.Register(etcdResolver) + + grpcFactory := generic.NewGpcGenericClientFactory(etcdResolver, grpc.WithTransportCredentials(insecure.NewCredentials())) + grpcFactory.Init() + + // gin http server + authFilter, err := config2.NewAuthFilter(appConf.AesTokenKey, appConf.IgnoreUrls) + if err != nil { + panic(err) + } + + server := gin.Default() + server.Use(authFilter.Filter) + + group := server.Group("/api/svc") + group.POST("/:svc/:method", func(c *gin.Context) { + svc := strs.UpperInitialLetter(c.Param("svc")) + method := strs.UpperInitialLetter(c.Param("method")) + if svc == "" || method == "" { + c.JSON(http.StatusBadRequest, resp.Error("svc not found")) + return + } + + ctx := context.Background() + grpcClient, err := grpcFactory.GetClient(ctx, svc) + if err != nil { + c.JSON(http.StatusForbidden, resp.Error(err.Error())) + return + } + body := make(map[string]interface{}) + err = c.BindJSON(&body) + if err != nil { + c.JSON(http.StatusBadRequest, resp.Error("parse request body error: "+err.Error())) + return + } + + res, err := grpcClient.InvokeUnaryJson(ctx, method, body) + if err != nil { + c.JSON(http.StatusInternalServerError, resp.Error(err.Error())) + return + } + // j, err := res.MarshalJSON() + // c.Render(http.StatusOK, RenderMarshaledJson{j}) + c.JSON(http.StatusOK, resp.Success(res)) + }) + + go func() { + err := server.Run(fmt.Sprintf(":%d", appConf.Port)) + if err != nil { + panic(err) + } + }() + + shutdown.Await() +} + +var jsonContentType = []string{"application/json; charset=utf-8"} + +type RenderMarshaledJson struct { + MarshaledJson []byte +} + +func (r RenderMarshaledJson) Render(writer http.ResponseWriter) error { + r.WriteContentType(writer) + _, err := writer.Write(r.MarshaledJson) + return err +} + +func (r RenderMarshaledJson) WriteContentType(w http.ResponseWriter) { + header := w.Header() + header["Content-Type"] = jsonContentType +} diff --git a/cmd/gateway_ws/config.toml b/cmd/gateway_ws/config.toml index 1cca101..69ce25d 100644 --- a/cmd/gateway_ws/config.toml +++ b/cmd/gateway_ws/config.toml @@ -1,5 +1,8 @@ [app] httpPort = 7001 +subjectCacheTopic = "wsgate:subject:" +subjectLrcExpiration = "10m" +subjectLrcCleanupInterval = "5m" [grpc] address = ":7010" diff --git a/cmd/gateway_ws/main.go b/cmd/gateway_ws/main.go index 955fbfd..3f296b4 100644 --- a/cmd/gateway_ws/main.go +++ b/cmd/gateway_ws/main.go @@ -1,6 +1,8 @@ package main import ( + "github.com/nats-io/nats.go" + "github.com/redis/go-redis/v9" clientv3 "go.etcd.io/etcd/client/v3" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" @@ -9,14 +11,23 @@ import ( "sonet/internal/gateway_ws/server" "sonet/internal/postal/logic" "sonet/pkg/config" + "sonet/pkg/grpc/client" "sonet/pkg/grpc/discovery" "sonet/pkg/grpc/generic" + "sonet/pkg/plugins/cache" + "sonet/pkg/plugins/mq" + "sonet/pkg/utils/conver" "sonet/pkg/utils/logger" "sonet/pkg/utils/shutdown" + "sync" ) type GatewayWsConfig struct { HttpPort int + + SubjectCacheTopic string + SubjectLrcExpiration string + SubjectLrcCleanupInterval string } // websocket server with postalService @@ -26,18 +37,26 @@ func main() { appConf := &GatewayWsConfig{} conf := config.LoadConfig(appConf, "cmd/gateway_ws") - // registry and run... + // registry and run etcdClient, err := clientv3.New(conf.Etcd) if err != nil { panic(err) } + sessionStore := &sync.Map{} + subjectStore, err := getSubjectMultiCache(appConf, conf.Redis, conf.Nats) + if err != nil { + panic(err) + } + clientFactory := client.NewGrpcDirectClientFactory(grpc.WithTransportCredentials(insecure.NewCredentials())) etcdResolver := discovery.NewResolver(etcdClient) resolver.Register(etcdResolver) grpcFactory := generic.NewGpcGenericClientFactory(etcdResolver, grpc.WithTransportCredentials(insecure.NewCredentials())) grpcFactory.Init() - connHandler := server.NewConnHandler(grpcFactory) + + postalAddr := discovery.MustRegisterAddress(conf.Grpc.Address) + connHandler := server.NewConnHandler(postalAddr, grpcFactory, sessionStore, subjectStore) httpServer := server.NewHttpServer(connHandler) go func() { err = httpServer.Run(appConf.HttpPort) @@ -46,9 +65,10 @@ func main() { } }() + // run postal server registry := discovery.NewRegister(etcdClient) shutdown.AddShutdownHook(registry.Stop) - postalServer := logic.NewPostalServer() + postalServer := logic.NewPostalServer(sessionStore, subjectStore, clientFactory) go func() { err = postalServer.Run(conf.Grpc, registry) if err != nil { @@ -56,5 +76,45 @@ func main() { } }() + // run postal cluster server + postalClusterServer := logic.NewPostalClusterServer(postalServer) + go func() { + err := postalClusterServer.Run(conf.Grpc.Address, config.GetGrpcOptions(conf.Grpc)...) + if err != nil { + panic(err) + } + }() + shutdown.Await() } + +func getSubjectMultiCache(appConf *GatewayWsConfig, redisOptions redis.Options, natsOptions nats.Options) (*cache.LocalRemoteCache, error) { + // initial cache... + rdb := redis.NewClient(&redisOptions) + subjectRedisCache := cache.NewRedisCache(appConf.SubjectCacheTopic, rdb) + shutdown.AddShutdownHook(func() { _ = rdb.Close() }) + + // nats mq + producer, err := mq.NewNatsProducer(natsOptions) + if err != nil { + return nil, err + } + shutdown.AddShutdownHook(func() { producer.Stop() }) + consumer, err := mq.NewNatsConsumer(natsOptions) + if err != nil { + return nil, err + } + shutdown.AddShutdownHook(func() { consumer.Stop() }) + + // 多级缓存 + subjectLrcOpts := cache.LocalRemoteCacheOptions{ + Topic: appConf.SubjectCacheTopic, + LocalExpiration: conver.MustParseDuration(appConf.SubjectLrcExpiration), + CleanupInterval: conver.MustParseDuration(appConf.SubjectLrcCleanupInterval), + Remote: subjectRedisCache, + Producer: producer, + Consumer: consumer, + } + subjectLrc, err := cache.NewLocalRemoteCache(subjectLrcOpts) + return subjectLrc, err +} diff --git a/cmd/mahjong/main.go b/cmd/mahjong/main.go new file mode 100644 index 0000000..b79dd79 --- /dev/null +++ b/cmd/mahjong/main.go @@ -0,0 +1,30 @@ +package main + +import ( + "fmt" + "net" +) + +func main() { + + //listen, err := net.Listen("tcp", ":7878") + //addr := listen.(*net.TCPListener).Addr() + //port := addr.(*net.TCPAddr).Port + //fmt.Println(err, listen, addr, port) + // listen.(*net.TCPListener).Addr().(*net.TCPAddr).Port + + addr, err := net.ResolveTCPAddr("tcp", "192.168.1.110:7788") + fmt.Println(err) + fmt.Println(addr.Port) + if addr.IP == nil { + fmt.Println("ip is nil") + } + fmt.Println(addr.IP.String()) + + //addrPort, err := netip.ParseAddrPort(":7788") + //fmt.Println(err) + //fmt.Println(addrPort.Port()) + //fmt.Println(addrPort.Addr()) + //fmt.Println(addrPort.String()) + +} diff --git a/go.mod b/go.mod index b071139..b129c00 100644 --- a/go.mod +++ b/go.mod @@ -7,12 +7,14 @@ require ( github.com/dsnet/golib/unitconv v1.0.2 github.com/gin-gonic/gin v1.9.1 github.com/golang/protobuf v1.5.3 + github.com/google/uuid v1.4.0 github.com/gorilla/websocket v1.5.1 github.com/jhump/protoreflect v1.15.4 github.com/nats-io/nats.go v1.31.0 + github.com/patrickmn/go-cache v2.1.0+incompatible github.com/redis/go-redis/v9 v9.4.0 github.com/sirupsen/logrus v1.9.3 - github.com/spf13/viper v1.18.2 + github.com/spf13/viper v1.17.0 go.etcd.io/etcd/client/v3 v3.5.11 google.golang.org/grpc v1.60.1 google.golang.org/protobuf v1.32.0 @@ -41,7 +43,7 @@ require ( github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect github.com/json-iterator/go v1.1.12 // indirect - github.com/klauspost/compress v1.17.0 // indirect + github.com/klauspost/compress v1.17.4 // indirect github.com/klauspost/cpuid/v2 v2.2.4 // indirect github.com/leodido/go-urn v1.2.4 // indirect github.com/magiconair/properties v1.8.7 // indirect @@ -49,7 +51,7 @@ require ( github.com/mitchellh/mapstructure v1.5.0 // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect github.com/modern-go/reflect2 v1.0.2 // indirect - github.com/nats-io/nkeys v0.4.6 // indirect + github.com/nats-io/nkeys v0.4.7 // indirect github.com/nats-io/nuid v1.0.1 // indirect github.com/pelletier/go-toml/v2 v2.1.0 // indirect github.com/sagikazarmark/locafero v0.4.0 // indirect @@ -67,11 +69,11 @@ require ( go.uber.org/multierr v1.9.0 // indirect go.uber.org/zap v1.21.0 // indirect golang.org/x/arch v0.3.0 // indirect - golang.org/x/crypto v0.16.0 // indirect + golang.org/x/crypto v0.18.0 // indirect golang.org/x/exp v0.0.0-20230905200255-921286631fa9 // indirect golang.org/x/net v0.19.0 // indirect golang.org/x/sync v0.5.0 // indirect - golang.org/x/sys v0.15.0 // indirect + golang.org/x/sys v0.16.0 // indirect golang.org/x/text v0.14.0 // indirect google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17 // indirect google.golang.org/genproto/googleapis/api v0.0.0-20231106174013-bbf56f31fb17 // indirect diff --git a/go.sum b/go.sum index f139422..41ac8e6 100644 --- a/go.sum +++ b/go.sum @@ -58,6 +58,8 @@ github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMyw github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= +github.com/google/uuid v1.4.0 h1:MtMxsa51/r9yyhkyLsVeVt0B+BGQZzpQiTQ4eHZ8bc4= +github.com/google/uuid v1.4.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/websocket v1.5.1 h1:gmztn0JnHVt9JZquRuzLw3g4wouNVzKL15iLr/zn/QY= github.com/gorilla/websocket v1.5.1/go.mod h1:x3kM2JMyaluk02fnUJpQuwD2dCS5NDG2ZHL0uE0tcaY= github.com/hashicorp/hcl v1.0.0 h1:0Anlzjpi4vEasTeNFn2mLJgTSwt0+6sfsiTG8qcWGx4= @@ -72,8 +74,8 @@ github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnr github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= -github.com/klauspost/compress v1.17.0 h1:Rnbp4K9EjcDuVuHtd0dgA4qNuv9yKDYKK1ulpJwgrqM= -github.com/klauspost/compress v1.17.0/go.mod h1:ntbaceVETuRiXiv4DpjP66DpAtAGkEQskQzEyD//IeE= +github.com/klauspost/compress v1.17.4 h1:Ej5ixsIri7BrIjBkRZLTo6ghwrEtHFk7ijlczPW4fZ4= +github.com/klauspost/compress v1.17.4/go.mod h1:/dCuZOvVtNoHsyb+cuJD3itjs3NbnF6KH9zAO4BDxPM= github.com/klauspost/cpuid/v2 v2.0.9/go.mod h1:FInQzS24/EEf25PyTYn52gqo7WaD8xa0213Md/qVLRg= github.com/klauspost/cpuid/v2 v2.2.4 h1:acbojRNwl3o09bUq+yDCtZFc1aiwaAAxtcn8YkZXnvk= github.com/klauspost/cpuid/v2 v2.2.4/go.mod h1:RVVoqg1df56z8g3pUjL/3lE5UfnlrJX8tyFgg4nqhuY= @@ -98,10 +100,12 @@ github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9G github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= github.com/nats-io/nats.go v1.31.0 h1:/WFBHEc/dOKBF6qf1TZhrdEfTmOZ5JzdJ+Y3m6Y/p7E= github.com/nats-io/nats.go v1.31.0/go.mod h1:di3Bm5MLsoB4Bx61CBTsxuarI36WbhAwOm8QrW39+i8= -github.com/nats-io/nkeys v0.4.6 h1:IzVe95ru2CT6ta874rt9saQRkWfe2nFj1NtvYSLqMzY= -github.com/nats-io/nkeys v0.4.6/go.mod h1:4DxZNzenSVd1cYQoAa8948QY3QDjrHfcfVADymtkpts= +github.com/nats-io/nkeys v0.4.7 h1:RwNJbbIdYCoClSDNY7QVKZlyb/wfT6ugvFCiKy6vDvI= +github.com/nats-io/nkeys v0.4.7/go.mod h1:kqXRgRDPlGy7nGaEDMuYzmiJCIAAWDK0IMBtDmGD0nc= github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw= github.com/nats-io/nuid v1.0.1/go.mod h1:19wcPz3Ph3q0Jbyiqsd0kePYG7A95tJPxeL+1OSON2c= +github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc= +github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ= github.com/pelletier/go-toml/v2 v2.1.0 h1:FnwAJ4oYMvbT/34k9zzHuZNrhlz48GB3/s6at6/MHO4= github.com/pelletier/go-toml/v2 v2.1.0/go.mod h1:tJU2Z3ZkXwnxa4DPO899bsyIoywizdUvyaeZurnPPDc= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= @@ -125,8 +129,8 @@ github.com/spf13/cast v1.6.0 h1:GEiTHELF+vaR5dhz3VqZfFSzZjYbgeKDpBxQVS4GYJ0= github.com/spf13/cast v1.6.0/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo= github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= -github.com/spf13/viper v1.18.2 h1:LUXCnvUvSM6FXAsj6nnfc8Q2tp1dIgUfY9Kc8GsSOiQ= -github.com/spf13/viper v1.18.2/go.mod h1:EKmWIqdnk5lOcmR72yw6hS+8OPYcwD0jteitLMVB+yk= +github.com/spf13/viper v1.17.0 h1:I5txKw7MJasPL/BrfkbA0Jyo/oELqVmux4pR/UxOMfI= +github.com/spf13/viper v1.17.0/go.mod h1:BmMMMLQXSbcHK6KAOiFLz0l5JHrU89OdIRHvsk0+yVI= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= @@ -169,8 +173,8 @@ golang.org/x/arch v0.3.0/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/crypto v0.16.0 h1:mMMrFzRSCF0GvB7Ne27XVtVAaXLrPmgPC7/v0tkwHaY= -golang.org/x/crypto v0.16.0/go.mod h1:gCAAfMLgwOJRpTjQ2zCCt2OcSfYMTeZVSRtQlPC7Nq4= +golang.org/x/crypto v0.18.0 h1:PGVlW0xEltQnzFZ55hkuX5+KLyrMYhHld1YHO4AKcdc= +golang.org/x/crypto v0.18.0/go.mod h1:R0j02AL6hcrfOiy9T4ZYp/rcWeMxM3L6QYxlOuEG1mg= golang.org/x/exp v0.0.0-20230905200255-921286631fa9 h1:GoHiUyI/Tp2nVkLI2mCxVkOjsbSXD66ic0XW0js0R9g= golang.org/x/exp v0.0.0-20230905200255-921286631fa9/go.mod h1:S2oDrQGGwySpoQPVqRShND87VCbxmc6bL1Yd2oYrm6k= golang.org/x/lint v0.0.0-20190930215403-16217165b5de/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= @@ -200,8 +204,8 @@ golang.org/x/sys v0.0.0-20210510120138-977fb7262007/go.mod h1:oPkhp1MJrh7nUepCBc golang.org/x/sys v0.0.0-20220704084225-05e143d24a9e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.15.0 h1:h48lPFYpsTvQJZF4EKyI4aLHaev3CxivZmv7yZig9pc= -golang.org/x/sys v0.15.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +golang.org/x/sys v0.16.0 h1:xWw16ngr6ZMtmxDyKyIgsE93KNKz5HKmMa3b8ALHidU= +golang.org/x/sys v0.16.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= diff --git a/internal/auth/data/auth_user_dao.go b/internal/auth/data/auth_user_dao.go new file mode 100644 index 0000000..c9de8ed --- /dev/null +++ b/internal/auth/data/auth_user_dao.go @@ -0,0 +1,66 @@ +package data + +import ( + "errors" + "gorm.io/gorm" + "strconv" +) + +type AuthUserDao struct { + db *gorm.DB +} + +func NewUserDao(db *gorm.DB) *AuthUserDao { + return &AuthUserDao{db: db} +} + +func (dao *AuthUserDao) Create(user *User) error { + tx := dao.db.Create(user) + return tx.Error +} + +func (dao *AuthUserDao) Updates(user *User) error { + return dao.db.Model(user).Updates(user).Error +} + +func (dao *AuthUserDao) FindByAccount(account string) (*User, error) { + user := &User{} + tx := dao.db.Model(user).Where("account=?", account).Take(user) + if tx.Error != nil { + // First、Last、Take 方法找不到记录时,ErrRecordNotFound + if errors.Is(tx.Error, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, tx.Error + } + return user, nil +} + +func (dao *AuthUserDao) FindByUid(uid string) (*User, error) { + user := &User{} + tx := dao.db.Model(user).Where("uid=?", uid).Take(user) + if tx.Error != nil { + // First、Last、Take 方法找不到记录时,ErrRecordNotFound + if errors.Is(tx.Error, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, tx.Error + } + return user, nil +} + +func (dao *AuthUserDao) NextUid() (string, error) { + user := &User{} + tx := dao.db.Model(user).Where("uid=(select MAX(uid) from `t_user`)").Take(user) + if tx.Error != nil { + if errors.Is(tx.Error, gorm.ErrRecordNotFound) { + return "10000", nil + } + return "", tx.Error + } + uid, err := strconv.Atoi(user.Uid) + if err != nil { + return "", err + } + return strconv.Itoa(uid + 1), nil +} diff --git a/internal/auth/data/t_user.go b/internal/auth/data/t_user.go new file mode 100644 index 0000000..5c3d22d --- /dev/null +++ b/internal/auth/data/t_user.go @@ -0,0 +1,18 @@ +package data + +import "time" + +// User 用户 +type User struct { + Uid string `gorm:"column:uid;primaryKey;" json:"uid"` // 用户id + Account string `gorm:"column:account" json:"account"` // 登录账号 + Username string `gorm:"column:username" json:"username"` // 用户名 + Password string `gorm:"column:password" json:"password"` // 密码 + Avatar string `gorm:"column:avatar" json:"avatar"` // 头像 + CreateAt *time.Time `gorm:"column:create_at" json:"createAt"` // 创建时间 + UpdateAt *time.Time `gorm:"column:update_at" json:"updateAt"` // 更新时间 +} + +func (User) TableName() string { + return "t_user" +} diff --git a/internal/auth/logic/auth_server.go b/internal/auth/logic/auth_server.go index 7432e93..b212ca4 100644 --- a/internal/auth/logic/auth_server.go +++ b/internal/auth/logic/auth_server.go @@ -2,22 +2,36 @@ package logic import ( "context" + "encoding/base64" "errors" + "github.com/bytedance/sonic" "google.golang.org/grpc" "google.golang.org/grpc/reflection" "net" "sonet/api/gen/auth" + "sonet/internal/auth/data" "sonet/pkg/config" "sonet/pkg/grpc/discovery" "sonet/pkg/utils/logger" + "sonet/pkg/utils/security" + "time" ) type AuthServer struct { auth.UnimplementedAuthServer + aesTokenKey []byte + userDao *data.AuthUserDao } -func NewAuthServer() *AuthServer { - return &AuthServer{} +func NewAuthServer(aesTokenKey string, userDao *data.AuthUserDao) *AuthServer { + keyBytes, err := base64.StdEncoding.DecodeString(aesTokenKey) + if err != nil { + panic(err) + } + return &AuthServer{ + aesTokenKey: keyBytes, + userDao: userDao, + } } func (s *AuthServer) Run(conf config.GrpcConfig, register *discovery.Register) (err error) { @@ -35,38 +49,105 @@ func (s *AuthServer) Run(conf config.GrpcConfig, register *discovery.Register) ( } // registry discovery - registerConf := conf.Register - if registerConf.Name == "" { - registerConf.Name = auth.Auth_ServiceDesc.ServiceName + reg := conf.Register + if reg.Name == "" { + reg.Name = auth.Auth_ServiceDesc.ServiceName } - if registerConf.Addr == "" { - registerConf.Addr, err = discovery.RegisterAddress(listen) + if reg.Addr == "" { + reg.Addr, err = discovery.RegisterAddress(conf.Address) if err != nil { return } } - if err = register.Register(registerConf); err != nil { + if err = register.Register(reg); err != nil { return } // run serve - logger.Infof("%s grpc server running %s\n", registerConf.Name, listen.Addr().String()) + logger.Infof("%s grpc server running %s\n", reg.Name, listen.Addr().String()) err = server.Serve(listen) return } func (s *AuthServer) Login(ctx context.Context, req *auth.ReqLogin) (*auth.ResLogin, error) { - return nil, nil + user, err := s.userDao.FindByAccount(req.Account) + if err != nil { + return nil, err + } + if user == nil { + // 不存在注册 + now := time.Now() + uid, err := s.userDao.NextUid() + if err != nil { + return nil, err + } + user = &data.User{ + Uid: uid, + Account: req.Account, + Username: req.Account, + Password: req.Password, + CreateAt: &now, + UpdateAt: &now, + } + err = s.userDao.Create(user) + if err != nil { + return nil, err + } + // return nil, errors.New("not found account " + req.Account) + } + // verify password login + if req.Password != user.Password { + return nil, errors.New("account or password error") + } + + // generate token + subject := &auth.Subject{Uid: user.Uid, Username: user.Username, Time: time.Now().UnixMilli()} + bytes, err := sonic.Marshal(subject) + if err != nil { + return nil, err + } + encode, err := security.EncryptAesCBC(bytes, s.aesTokenKey) + if err != nil { + return nil, err + } + token := base64.URLEncoding.EncodeToString(encode) + + res := &auth.ResLogin{Token: token, Subject: subject} + return res, nil } func (s *AuthServer) Verify(ctx context.Context, req *auth.ReqVerify) (*auth.Subject, error) { - return &auth.Subject{Username: "sunshine"}, errors.New("test error...") + bytes, err := base64.URLEncoding.DecodeString(req.Token) + if err != nil { + return nil, errors.New("token decode fail: " + err.Error()) + } + decode, err := security.DecryptAesCBC(bytes, s.aesTokenKey) + if err != nil { + return nil, errors.New("invalidate token: " + err.Error()) + } + subject := &auth.Subject{} + err = sonic.Unmarshal(decode, subject) + if err != nil { + return nil, errors.New("token payload decode fail: " + err.Error()) + } + return subject, nil } func (s *AuthServer) Registry(ctx context.Context, req *auth.ReqRegistry) (*auth.Subject, error) { - return nil, nil + return nil, errors.New("can not registry") } -func (s *AuthServer) FindByUid(ctx context.Context, req *auth.ReqFindByUid) (*auth.Subject, error) { - return nil, nil +func (s *AuthServer) FindByUid(ctx context.Context, req *auth.ReqFindByUid) (subject *auth.Subject, err error) { + sub, err := s.userDao.FindByUid(req.Uid) + if err != nil || sub == nil { + return + } + + subject = &auth.Subject{ + Uid: sub.Uid, + Username: sub.Username, + Time: time.Now().UnixMilli(), + Extra: map[string]string{"avatar": sub.Avatar}, + } + return } diff --git a/internal/chat/logic/chat_server.go b/internal/chat/logic/chat_server.go index b686da9..30cd859 100644 --- a/internal/chat/logic/chat_server.go +++ b/internal/chat/logic/chat_server.go @@ -2,23 +2,80 @@ package logic import ( "context" + "google.golang.org/grpc" + "google.golang.org/grpc/reflection" "google.golang.org/protobuf/types/known/emptypb" + "net" "sonet/api/gen/chat" - "sonet/api/gen/postal" + "sonet/pkg/config" + "sonet/pkg/grpc/discovery" + "sonet/pkg/protocol/deliver" + "sonet/pkg/protocol/session" + "sonet/pkg/utils/logger" ) type ChatServer struct { chat.UnimplementedChatServer - postalCli postal.PostalClient + deliver *deliver.Deliver } -func NewChatServer() *ChatServer { - return &ChatServer{} +func NewChatServer(deliver *deliver.Deliver) *ChatServer { + return &ChatServer{ + deliver: deliver, + } } -func (s *ChatServer) Send(ctx context.Context, send *chat.ReqSend) (*chat.ResSend, error) { +func (s *ChatServer) Run(conf config.GrpcConfig, register *discovery.Register) (err error) { + server := grpc.NewServer( + config.GetGrpcOptions(conf)..., + ) + if !conf.NoReflection { + // 注册反射服务 + reflection.Register(server) + } + chat.RegisterChatServer(server, s) + listen, err := net.Listen("tcp", conf.Address) + if err != nil { + return + } - return nil, nil + // registry discovery + reg := conf.Register + if reg.Name == "" { + reg.Name = chat.Chat_ServiceDesc.ServiceName + } + if reg.Addr == "" { + reg.Addr, err = discovery.RegisterAddress(conf.Address) + if err != nil { + return + } + } + if err = register.Register(reg); err != nil { + return + } + + // run serve + logger.Infof("%s grpc server running %s\n", reg.Name, listen.Addr().String()) + err = server.Serve(listen) + return +} + +func (s *ChatServer) Send(ctx context.Context, req *chat.ReqSend) (*chat.ResSend, error) { + subject, err := session.GetSubject(ctx) + if err != nil { + return nil, err + } + + // TODO 存储消息,ack 队列重发(可异步发送),asc超时第二次发送时同步判断是否不在线? + + // 投递消息 + message := &chat.ChatMessage{Sender: subject.Uid, Content: req.Content} + _, err = s.deliver.Deliver(ctx, message, req.Receiver) + if err != nil { + return nil, err + } + + return &chat.ResSend{}, nil } func (s *ChatServer) RoomSend(ctx context.Context, send *chat.ReqRoomSend) (*chat.ResSend, error) { diff --git a/internal/gateway_http/config/auth_filter.go b/internal/gateway_http/config/auth_filter.go new file mode 100644 index 0000000..ca048fe --- /dev/null +++ b/internal/gateway_http/config/auth_filter.go @@ -0,0 +1,55 @@ +package config + +import ( + "encoding/base64" + "github.com/gin-gonic/gin" + "net/http" + "sonet/pkg/protocol/authorize" + "sonet/pkg/utils/resp" + "strings" +) + +var SubjectKey = "session:subject" + +type AuthFilter struct { + ignoreUrls []string + aesTokenKey []byte +} + +func NewAuthFilter(aesTokenKey string, ignoreUrls []string) (*AuthFilter, error) { + keyBytes, err := base64.StdEncoding.DecodeString(aesTokenKey) + if err != nil { + panic(err) + } + return &AuthFilter{ + ignoreUrls: ignoreUrls, + aesTokenKey: keyBytes, + }, nil +} + +func (f *AuthFilter) Filter(c *gin.Context) { + path := c.Request.URL.Path + + // ignore auth urls + for _, ignoreUrl := range f.ignoreUrls { + if strings.EqualFold(ignoreUrl, path) { + return + } + } + + token := c.GetHeader("Authorization") + if token == "" { + c.JSON(http.StatusUnauthorized, resp.Fail("请登录后进行操作")) + c.Abort() + return + } + + subject, err := authorize.Verify(f.aesTokenKey, token) + if err != nil { + c.JSON(http.StatusUnauthorized, resp.Fail("认证失败,请重新登录")) + c.Abort() + return + } + + c.Set(SubjectKey, subject) +} diff --git a/internal/gateway_http/logic/http_server.go b/internal/gateway_http/logic/http_server.go new file mode 100644 index 0000000..4c79103 --- /dev/null +++ b/internal/gateway_http/logic/http_server.go @@ -0,0 +1 @@ +package logic diff --git a/internal/gateway_ws/server/ws_server.go b/internal/gateway_ws/server/ws_server.go index 794401c..2287b17 100644 --- a/internal/gateway_ws/server/ws_server.go +++ b/internal/gateway_ws/server/ws_server.go @@ -3,20 +3,35 @@ package server import ( "context" "github.com/gorilla/websocket" + "google.golang.org/protobuf/proto" "runtime/debug" + "sonet/api/gen/auth" "sonet/internal/gateway_ws/session" "sonet/pkg/grpc/generic" + "sonet/pkg/plugins/cache" "sonet/pkg/protocol" + session2 "sonet/pkg/protocol/session" "sonet/pkg/utils/logger" + "sync" + "time" ) type ConnHandler struct { - grpcFactory *generic.GrpcGenericClientFactory + postalServerAddress string + grpcFactory *generic.GrpcGenericClientFactory + sessionStore *sync.Map // 当前连接用户,内存缓存 + subjectStore cache.MultiLevelCache } -func NewConnHandler(grpcFactory *generic.GrpcGenericClientFactory) *ConnHandler { +func NewConnHandler(postalServerAddress string, + grpcFactory *generic.GrpcGenericClientFactory, + sessionStore *sync.Map, + subjectStore cache.MultiLevelCache) *ConnHandler { return &ConnHandler{ - grpcFactory: grpcFactory, + grpcFactory: grpcFactory, + postalServerAddress: postalServerAddress, + sessionStore: sessionStore, + subjectStore: subjectStore, } } @@ -24,8 +39,6 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { client := session.NewNetClient(conn, ReadDeadline, WriteDeadline) defer func() { - // TODO 连接关闭,mq发送关闭事件 - // 捕获其他错误 if r := recover(); r != nil { logger.Error("NetClient recover error: ", r) @@ -34,7 +47,20 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { } }() - defer client.Close() + defer func() { + client.Close() + // TODO 连接关闭,mq发送关闭事件 + // 删除连接 + if client.Subject != nil { + c.sessionStore.Delete(client.Subject.Uid) + err := c.subjectStore.Del(context.Background(), client.Subject.Uid) + if err != nil { + logger.Error("del cluster subject store uid error: ", client.Subject.Uid) + } + logger.Info("subject offline: ", client.Subject.Uid) + client.Subject = nil + } + }() for { ignore, message, err := client.ReadMessage() @@ -58,6 +84,11 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { } header := payload.Header + if client.Subject == nil && !(header.Svc == "Auth" && header.Target == "Verify") { + writeError(client, header.SeqId, session2.UnauthorizedRequestError.Error()) + continue + } + // grpc generic call ctx := context.Background() grpcClient, err := c.grpcFactory.GetClient(ctx, header.Svc) @@ -65,6 +96,10 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { logger.Error("get grpc generic client error: ", err) continue } + // put session + if client.Subject != nil { + ctx = session2.PutSubject(ctx, session2.NewRpcSubject(client.Subject.Uid)) + } resp, err := grpcClient.InvokeUnary(ctx, header.Target, payload.Body) if err != nil { // todo 提取 err: fmt.Sprintf("rpc error: code = %s desc = %s", s.Code(), s.Message()) @@ -74,13 +109,25 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { } // write response - header.Type = protocol.TypeResponse - header.Target = resp.XXX_MessageName() payload.Body, err = resp.Marshal() if err != nil { logger.Error("generic call response marshal error: ", err) + writeError(client, header.SeqId, "server error") continue } + + // auth verify success + if header.Svc == "Auth" && header.Target == "Verify" { + err = c.extractAuthVerify(client, payload.Body) + if err != nil { + logger.Error("extractAuthVerify error: ", err) + writeError(client, header.SeqId, "server error") + continue + } + } + + header.Type = protocol.TypeResponse + header.Target = resp.XXX_MessageName() resMessage, err := protocol.EncodeSo(payload) if err != nil { logger.Error("grpc generic call error: ", err) @@ -88,7 +135,38 @@ func (c *ConnHandler) handleConn(conn *websocket.Conn) { } client.MustWrite(resMessage) } +} +func (c *ConnHandler) extractAuthVerify(netClient *session.NetClient, resp []byte) (err error) { + authSubject := &auth.Subject{} + err = proto.Unmarshal(resp, authSubject) + if err != nil { + logger.Error("grpc generic call error: ", err) + return + } + // 设置连接身份信息 + subject := &session.Subject{ + Uid: authSubject.Uid, + Online: 1, + Time: time.Now().UnixMilli(), + Gate: c.postalServerAddress, + } + // 存储session + err = c.subjectStore.Set(context.Background(), subject.Uid, subject) // store to redis cluster cache + if err != nil { + return + } + + // 关闭旧的链接 + oldOnline, ok := c.sessionStore.Load(subject.Uid) + if ok { + oldOnline.(*session.NetSubject).Client.Close() // TODO nats offline / force load subject target cluster call offline + logger.Infof("close subject old conn: %s\n", subject.Uid) + } + netClient.Subject = session.NewNetSubject(subject.Uid, netClient) + c.sessionStore.Store(netClient.Subject.Uid, netClient.Subject) + logger.Infof("subject online: %s\n", subject.Uid) + return } func writeError(client *session.NetClient, seqId int32, errMsg string) { diff --git a/internal/gateway_ws/session/net_account.go b/internal/gateway_ws/session/net_account.go deleted file mode 100644 index a081160..0000000 --- a/internal/gateway_ws/session/net_account.go +++ /dev/null @@ -1,15 +0,0 @@ -package session - -import ( - "sync" -) - -// NetAccount 已认证的长连接用户 -type NetAccount struct { - Id int64 - UserName string - Avatar string - Client *NetClient - - Lock *sync.Mutex -} diff --git a/internal/gateway_ws/session/net_client.go b/internal/gateway_ws/session/net_client.go index 301ec49..82e748d 100644 --- a/internal/gateway_ws/session/net_client.go +++ b/internal/gateway_ws/session/net_client.go @@ -13,34 +13,41 @@ import ( type NetClient struct { Conn *websocket.Conn writeDeadline, readDeadline time.Duration - Online *atomic.Bool - Account *NetAccount - WriteLock *sync.Mutex + writeLock *sync.Mutex + Subject *NetSubject + closed *atomic.Bool } func NewNetClient(conn *websocket.Conn, readDeadline, writeDeadline time.Duration) *NetClient { - online := &atomic.Bool{} - online.Store(true) - return &NetClient{ Conn: conn, readDeadline: readDeadline, writeDeadline: writeDeadline, - WriteLock: new(sync.Mutex), - Online: online, + writeLock: new(sync.Mutex), + closed: &atomic.Bool{}, } } func (c *NetClient) Close() { + c.writeLock.Lock() + defer c.writeLock.Unlock() + if c.closed.Load() { + return + } + c.closed.Store(true) err := c.Conn.Close() if err != nil { logger.Error("NetClient Close error: ", err) } } +func (c *NetClient) IsClosed() bool { + return c.closed.Load() +} + func (c *NetClient) Write(bytes []byte) (err error) { - c.WriteLock.Lock() - defer c.WriteLock.Unlock() + c.writeLock.Lock() + defer c.writeLock.Unlock() err = c.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err != nil { err = fmt.Errorf("NetClient SetWriteDeadline error: %s", err.Error()) diff --git a/internal/gateway_ws/session/net_subject.go b/internal/gateway_ws/session/net_subject.go new file mode 100644 index 0000000..0338141 --- /dev/null +++ b/internal/gateway_ws/session/net_subject.go @@ -0,0 +1,15 @@ +package session + +// NetSubject 已认证的长连接用户 +type NetSubject struct { + Uid string + + Client *NetClient // join device +} + +func NewNetSubject(uid string, client *NetClient) *NetSubject { + return &NetSubject{ + Uid: uid, + Client: client, + } +} diff --git a/internal/postal/logic/postal_cluster_server.go b/internal/postal/logic/postal_cluster_server.go new file mode 100644 index 0000000..437c429 --- /dev/null +++ b/internal/postal/logic/postal_cluster_server.go @@ -0,0 +1,77 @@ +package logic + +import ( + "context" + "errors" + "fmt" + "google.golang.org/grpc" + "net" + "sonet/api/gen/postal" + "sonet/pkg/utils/logger" +) + +// PostalClusterPortOffset 相对与 postal grpc server 接口偏移量 +const PostalClusterPortOffset = 1000 + +// PostalAddr2Cluster 根据 postal server 地址得到 postal cluster server 地址 +func PostalAddr2Cluster(postalServerAddr string) (string, error) { + addr, err := net.ResolveTCPAddr("tcp", postalServerAddr) + if err != nil { + return "", err + } + if addr.IP == nil { + return fmt.Sprintf(":%d", addr.Port+PostalClusterPortOffset), nil + } + return fmt.Sprintf("%s:%d", addr.IP.String(), addr.Port+PostalClusterPortOffset), nil +} + +var ( + methodRedirectDeliver int32 = 10 + methodRedirectDeliverBatch int32 = 11 + methodRedirectDeliverGroup int32 = 12 +) + +// PostalClusterServer 用于postal server集群间, 消息转发通信 +type PostalClusterServer struct { + postal.UnimplementedPostalClusterServer + postalServer postal.PostalServer // 当前节点 postal server +} + +func NewPostalClusterServer(postalServer postal.PostalServer) *PostalClusterServer { + return &PostalClusterServer{ + postalServer: postalServer, + } +} + +// Run 运行在同进程 postalServer 端口+1000 +func (s *PostalClusterServer) Run(postalAddr string, opts ...grpc.ServerOption) (err error) { + server := grpc.NewServer(opts...) + postal.RegisterPostalClusterServer(server, s) + // 解析端口 + address, err := PostalAddr2Cluster(postalAddr) + if err != nil { + return + } + listen, err := net.Listen("tcp", address) + if err != nil { + return + } + + // run serve + logger.Infof("%s grpc server running %s\n", postal.PostalCluster_ServiceDesc.ServiceName, listen.Addr().String()) + err = server.Serve(listen) + return +} + +func (s *PostalClusterServer) Redirect(ctx context.Context, req *postal.ReqRedirect) (*postal.ResDeliver, error) { + req.Ttl -= 1 + switch req.RedirectMethod { + case methodRedirectDeliver: + return s.postalServer.Deliver(ctx, req.Deliver) + case methodRedirectDeliverBatch: + return s.postalServer.DeliverBatch(ctx, req.DeliverBatch) + case methodRedirectDeliverGroup: + return s.postalServer.DeliverGroup(ctx, req.DeliverGroup) + } + return nil, errors.New(fmt.Sprintf("not found redirect type %d", req.RedirectMethod)) +} diff --git a/internal/postal/logic/postal_server.go b/internal/postal/logic/postal_server.go index 5f87e55..4074440 100644 --- a/internal/postal/logic/postal_server.go +++ b/internal/postal/logic/postal_server.go @@ -7,20 +7,35 @@ import ( "google.golang.org/protobuf/types/known/emptypb" "net" "sonet/api/gen/postal" + "sonet/internal/gateway_ws/session" "sonet/pkg/config" + "sonet/pkg/grpc/client" "sonet/pkg/grpc/discovery" + "sonet/pkg/plugins/cache" + "sonet/pkg/protocol" "sonet/pkg/utils/logger" + "sync" ) type PostalServer struct { postal.UnimplementedPostalServer + broadcastAddress string // postal server 在注册中心注册的地址, 其他服务可直连访问 + sessionStore *sync.Map // 在线用户conn存储 + subjectStore cache.MultiLevelCache // online subject address cache, ws gateway集群所有在线用户存储 + clientFactory *client.GrpcDirectClientFactory } -func NewPostalServer() *PostalServer { - return &PostalServer{} +func NewPostalServer(sessionStore *sync.Map, + subjectStore cache.MultiLevelCache, + clientFactory *client.GrpcDirectClientFactory) *PostalServer { + return &PostalServer{ + sessionStore: sessionStore, + subjectStore: subjectStore, + clientFactory: clientFactory, + } } -func (p *PostalServer) Run(conf config.GrpcConfig, postalRegister *discovery.Register) (err error) { +func (s *PostalServer) Run(conf config.GrpcConfig, postalRegister *discovery.Register) (err error) { server := grpc.NewServer( config.GetGrpcOptions(conf)..., ) @@ -28,57 +43,198 @@ func (p *PostalServer) Run(conf config.GrpcConfig, postalRegister *discovery.Reg // 注册反射服务 reflection.Register(server) } - postal.RegisterPostalServer(server, p) + postal.RegisterPostalServer(server, s) listen, err := net.Listen("tcp", conf.Address) if err != nil { return } // registry discovery - regConf := conf.Register - if regConf.Name == "" { - regConf.Name = postal.Postal_ServiceDesc.ServiceName + reg := conf.Register + if reg.Name == "" { + reg.Name = postal.Postal_ServiceDesc.ServiceName } - if regConf.Addr == "" { - regConf.Addr, err = discovery.RegisterAddress(listen) + if reg.Addr == "" { + reg.Addr, err = discovery.RegisterAddress(conf.Address) if err != nil { return } } - if err = postalRegister.Register(regConf); err != nil { + if err = postalRegister.Register(reg); err != nil { return } + // 其他服务直连地址 + s.broadcastAddress = reg.Addr // run serve - logger.Infof("%s grpc server running %s\n", regConf.Name, listen.Addr().String()) + logger.Infof("%s grpc server running %s\n", reg.Name, listen.Addr().String()) err = server.Serve(listen) return } -func (p *PostalServer) Deliver(ctx context.Context, deliver *postal.ReqDeliver) (*postal.ResDeliver, error) { - return nil, nil +func (s *PostalServer) deliverMessage(msg *postal.Message, netSubject *session.NetSubject) error { + header := &protocol.Header{ + Magic: protocol.Magic, + Type: protocol.TypeNotice, + UrlType: 1, + SerializeType: 1, + Svc: msg.Svc, + Target: msg.Msg, + } + payload := &protocol.Payload{Header: header, Body: msg.Body} + bytes, err := protocol.EncodeSo(payload) + if err != nil { + logger.Error("req deliver encode notice error: ", err) + return err + } + return netSubject.Client.Write(bytes) +} + +// receiverGates 找receiver在集群内哪些其他节点 +func (s *PostalServer) clusterGateReceivers(ctx context.Context, receivers []string) (gateReceivers map[string][]string, offline []string) { + gateReceivers = make(map[string][]string) + for _, receiver := range receivers { + subject := &session.Subject{} + err := s.subjectStore.Load(ctx, receiver, subject) + if err != nil { + offline = append(offline, receiver) + if err != cache.NotExists { + logger.Error("load from subject store error: ", err) + offline = append(offline, receiver) + } + continue + } + + if subject.Gate == s.broadcastAddress { + offline = append(offline, receiver) + // 清除失效缓存 + err := s.subjectStore.Del(ctx, receiver) + if err != nil { + logger.Error("del subject store error: ", receiver, err) + } + continue + } + + // put receiver gate addr + gateReceivers[subject.Gate] = append(gateReceivers[subject.Gate], receiver) + } + return } +func (s *PostalServer) Deliver(ctx context.Context, req *postal.ReqDeliver) (*postal.ResDeliver, error) { + val, ok := s.sessionStore.Load(req.Receiver) + + if ok { + receiver := val.(*session.NetSubject) + + err := s.deliverMessage(req.Msg, receiver) + if err != nil { + return nil, err + } + return &postal.ResDeliver{Ok: true}, nil + } + + // 用户连接不在当前gateway + gateReceivers, offline := s.clusterGateReceivers(ctx, []string{req.Receiver}) + if offline != nil && len(offline) > 0 { + return &postal.ResDeliver{Code: int32(postal.DeliverResult_ReceiverOffline.Number())}, nil + } + + for postalAddr := range gateReceivers { + // gateway集群中转发消息 + clusterAddr, err := PostalAddr2Cluster(postalAddr) + if err != nil { + logger.Errorf("parse postal server addr error: %s", postalAddr, err) + continue + } + conn, err := s.clientFactory.GetConn(context.Background(), clusterAddr) + if err != nil { + logger.Error("get postal cluster conn error: ", err) + return nil, err + } + clusterClient := postal.NewPostalClusterClient(conn) + + reqRedirect := &postal.ReqRedirect{ + Ttl: 3, // TODO 转发n次就丢弃 + RedirectMethod: methodRedirectDeliver, + Deliver: req, + } + resp, err := clusterClient.Redirect(ctx, reqRedirect) + return resp, err + } -func (p *PostalServer) DeliverBatch(ctx context.Context, batch *postal.ReqDeliverBatch) (*postal.ResDeliver, error) { return nil, nil } -func (p *PostalServer) DeliverGroup(ctx context.Context, group *postal.ReqDeliverGroup) (*postal.ResDeliver, error) { +func (s *PostalServer) DeliverBatch(ctx context.Context, req *postal.ReqDeliverBatch) (*postal.ResDeliver, error) { + var redirectReceivers []string + for _, receiverId := range req.Receivers { + val, ok := s.sessionStore.Load(receiverId) + if !ok { + redirectReceivers = append(redirectReceivers, receiverId) + continue + } + receiver := val.(*session.NetSubject) + err := s.deliverMessage(req.Msg, receiver) + if err != nil { + logger.Errorf("deliver to %s error: ", receiver, err) + } + } + + if len(redirectReceivers) == 0 { + return &postal.ResDeliver{Ok: true}, nil + } + + gateReceivers, offline := s.clusterGateReceivers(ctx, redirectReceivers) + if len(offline) > 0 { + logger.Warning("offline redirect receivers: ", offline) + } + + if len(gateReceivers) > 0 { + for postalAddr, receivers := range gateReceivers { + // gateway集群中转发消息 + clusterAddr, err := PostalAddr2Cluster(postalAddr) + if err != nil { + logger.Errorf("parse postal server addr error: %s", postalAddr, err) + continue + } + conn, err := s.clientFactory.GetConn(ctx, clusterAddr) + if err != nil { + logger.Error("get postal cluster conn error: ", err) + continue + } + clusterClient := postal.NewPostalClusterClient(conn) + + req.Receivers = receivers + reqRedirect := &postal.ReqRedirect{ + Ttl: 3, // TODO 转发n次就丢弃 + RedirectMethod: methodRedirectDeliverBatch, + DeliverBatch: req, + } + _, err = clusterClient.Redirect(ctx, reqRedirect) + if err != nil { + logger.Error("redirect batch error: ", err) + } + } + } + return &postal.ResDeliver{Ok: true}, nil +} + +func (s *PostalServer) DeliverGroup(ctx context.Context, group *postal.ReqDeliverGroup) (*postal.ResDeliver, error) { return nil, nil } -func (p *PostalServer) GroupCreate(ctx context.Context, create *postal.ReqGroupCreate) (*postal.ResGroupCreate, error) { +func (s *PostalServer) GroupCreate(ctx context.Context, create *postal.ReqGroupCreate) (*postal.ResGroupCreate, error) { return nil, nil } -func (p *PostalServer) GroupJoin(ctx context.Context, join *postal.ReqGroupJoin) (*emptypb.Empty, error) { +func (s *PostalServer) GroupJoin(ctx context.Context, join *postal.ReqGroupJoin) (*emptypb.Empty, error) { return nil, nil } -func (p *PostalServer) GroupLeave(ctx context.Context, leave *postal.ReqGroupLeave) (*emptypb.Empty, error) { +func (s *PostalServer) GroupLeave(ctx context.Context, leave *postal.ReqGroupLeave) (*emptypb.Empty, error) { return nil, nil } -func (p *PostalServer) GroupDissolve(ctx context.Context, dissolve *postal.ReqGroupDissolve) (*emptypb.Empty, error) { +func (s *PostalServer) GroupDissolve(ctx context.Context, dissolve *postal.ReqGroupDissolve) (*emptypb.Empty, error) { return nil, nil } diff --git a/pkg/config/grpc_options.go b/pkg/config/grpc_options.go index bdfb606..6850c91 100644 --- a/pkg/config/grpc_options.go +++ b/pkg/config/grpc_options.go @@ -18,6 +18,5 @@ func GetGrpcOptions(config GrpcConfig) (opts []grpc.ServerOption) { if config.WriteBufferSize != "" { opts = append(opts, grpc.MaxSendMsgSize(conver.MustParseDataUnitInt(config.WriteBufferSize))) } - return } diff --git a/pkg/deliver/deliver.go b/pkg/deliver/deliver.go deleted file mode 100644 index 9eae15d..0000000 --- a/pkg/deliver/deliver.go +++ /dev/null @@ -1 +0,0 @@ -package deliver diff --git a/pkg/grpc/client/direct_client_factory.go b/pkg/grpc/client/direct_client_factory.go new file mode 100644 index 0000000..6f6b2ce --- /dev/null +++ b/pkg/grpc/client/direct_client_factory.go @@ -0,0 +1,41 @@ +package client + +import ( + "context" + "google.golang.org/grpc" + "sync" +) + +type GrpcDirectClientFactory struct { + defaultOpts []grpc.DialOption + clientCache *sync.Map +} + +func NewGrpcDirectClientFactory(defaultOpts ...grpc.DialOption) *GrpcDirectClientFactory { + return &GrpcDirectClientFactory{ + defaultOpts: defaultOpts, + clientCache: &sync.Map{}, + } +} + +func (f *GrpcDirectClientFactory) NewConn(ctx context.Context, addr string, opts ...grpc.DialOption) (*grpc.ClientConn, error) { + dialOpts := make([]grpc.DialOption, 0, len(f.defaultOpts)+len(opts)) + dialOpts = append(dialOpts, f.defaultOpts...) + dialOpts = append(dialOpts, opts...) + + return grpc.DialContext(ctx, addr, dialOpts...) +} + +func (f *GrpcDirectClientFactory) GetConn(ctx context.Context, addr string, opts ...grpc.DialOption) (conn *grpc.ClientConn, err error) { + val, ok := f.clientCache.Load(addr) + if ok { + conn = val.(*grpc.ClientConn) + return + } + conn, err = f.NewConn(ctx, addr, opts...) + if err != nil { + return + } + f.clientCache.Store(addr, conn) + return +} diff --git a/pkg/grpc/discovery/register.go b/pkg/grpc/discovery/register.go index 531337d..e52a618 100644 --- a/pkg/grpc/discovery/register.go +++ b/pkg/grpc/discovery/register.go @@ -16,22 +16,31 @@ import ( var DefaultRegisterTTL int64 = 10 -func RegisterAddress(listener net.Listener) (addr string, err error) { - port := listener.Addr().(*net.TCPAddr).Port - ipv4, err := nets.GetHostIpv4() +func RegisterAddress(addr string) (fullAddr string, err error) { + tcpAddr, err := net.ResolveTCPAddr("tcp", addr) if err != nil { - return + panic(err) } - addr = fmt.Sprintf("%s:%d", ipv4, port) + var ip string + if tcpAddr.IP != nil { + ip = tcpAddr.IP.String() + } else { + ip, err = nets.GetHostIpv4() + if err != nil { + return + } + } + + fullAddr = fmt.Sprintf("%s:%d", ip, tcpAddr.Port) return } -func MustGetRegisterAddr(listener net.Listener) (addr string) { - port := listener.Addr().(*net.TCPAddr).Port - ipv4, err := nets.GetHostIpv4() + +func MustRegisterAddress(addr string) (fullAddr string) { + var err error + fullAddr, err = RegisterAddress(addr) if err != nil { panic(err) } - addr = fmt.Sprintf("%s:%s", ipv4, port) return } diff --git a/pkg/grpc/generic/generic_client.go b/pkg/grpc/generic/generic_client.go index 2875290..fceecf0 100644 --- a/pkg/grpc/generic/generic_client.go +++ b/pkg/grpc/generic/generic_client.go @@ -3,6 +3,8 @@ package generic import ( "context" "fmt" + "github.com/bytedance/sonic" + "github.com/golang/protobuf/proto" "github.com/jhump/protoreflect/desc" "github.com/jhump/protoreflect/dynamic" "github.com/jhump/protoreflect/dynamic/grpcdynamic" @@ -79,7 +81,37 @@ func (c *GrpcGenericClient) InvokeUnary(ctx context.Context, method string, reqB return } - res, err := caller.Stub.InvokeRpc(ctx, caller.Mtd, reqMessage, opts...) + return c.invokeUnary0(ctx, method, reqMessage, opts...) +} + +func (c *GrpcGenericClient) InvokeUnaryJson(ctx context.Context, method string, json map[string]interface{}, opts ...grpc.CallOption) (resp *dynamic.Message, err error) { + // cache method desc + caller, err := c.getMethodCaller(method) + if err != nil { + return + } + + reqMessage := caller.MsgFactory.NewMessage(caller.Mtd.GetInputType()) + jsonBytes, err := sonic.Marshal(json) + if err != nil { + return + } + if err = reqMessage.(*dynamic.Message).UnmarshalJSON(jsonBytes); err != nil { + err = fmt.Errorf("unmarshal req bytes error: %s", err.Error()) + return + } + + return c.invokeUnary0(ctx, method, reqMessage, opts...) +} + +func (c *GrpcGenericClient) invokeUnary0(ctx context.Context, method string, request proto.Message, opts ...grpc.CallOption) (resp *dynamic.Message, err error) { + // cache method desc + caller, err := c.getMethodCaller(method) + if err != nil { + return + } + + res, err := caller.Stub.InvokeRpc(ctx, caller.Mtd, request, opts...) if err != nil { return } diff --git a/pkg/grpc/generic/generic_client_factory.go b/pkg/grpc/generic/generic_client_factory.go index 5caf4d3..85f8f53 100644 --- a/pkg/grpc/generic/generic_client_factory.go +++ b/pkg/grpc/generic/generic_client_factory.go @@ -2,10 +2,9 @@ package generic import ( "context" - "errors" "fmt" "google.golang.org/grpc" - "runtime/debug" + "google.golang.org/grpc/resolver" "sonet/pkg/grpc/discovery" "sync" ) @@ -24,17 +23,11 @@ func NewGpcGenericClientFactory(resolver *discovery.Resolver, defaultOpts ...grp } func (f *GrpcGenericClientFactory) Init() { + resolver.Register(f.resolver) f.clientCache = &sync.Map{} } func (f *GrpcGenericClientFactory) NewClient(ctx context.Context, serviceName string, opts ...grpc.DialOption) (client *GrpcGenericClient, err error) { - defer func() { - if r := recover(); r != nil { - fmt.Println(r) - fmt.Println(string(debug.Stack())) - err = errors.New("errrrr") - } - }() addr := fmt.Sprintf("%s:///%s", f.resolver.Scheme(), serviceName) dialOpts := make([]grpc.DialOption, 0, len(f.defaultOpts)+len(opts)) dialOpts = append(dialOpts, f.defaultOpts...) diff --git a/pkg/plugins/cache/cache.go b/pkg/plugins/cache/cache.go new file mode 100644 index 0000000..720d8b2 --- /dev/null +++ b/pkg/plugins/cache/cache.go @@ -0,0 +1,46 @@ +package cache + +import ( + "context" + "encoding" +) + +const NotExists = CacheError("cache: not exists") + +type CacheError string + +func (e CacheError) Error() string { return string(e) } + +type TextSerializable interface { + encoding.TextMarshaler + encoding.TextUnmarshaler +} + +type BinarySerializable interface { + encoding.BinaryMarshaler + encoding.BinaryUnmarshaler +} + +// Cache TODO random Lock(key, timeout), Unlock(key, random) +type Cache interface { + Set(ctx context.Context, key string, value any) error + + Load(ctx context.Context, key string, target any) error + + Del(ctx context.Context, keys ...string) error +} + +type MultiLevelCache interface { + Cache + + ForceLoad(ctx context.Context, key string, target any) error +} + +// TODO CAS set value, atomic increment/decrement, 修改字段log复现 +// CasSet 对比当前版本 +func CasSet[V any](version int, value any) { + +} + +type Reply interface { +} diff --git a/pkg/plugins/cache/multi_cache.go b/pkg/plugins/cache/multi_cache.go new file mode 100644 index 0000000..69e0e6e --- /dev/null +++ b/pkg/plugins/cache/multi_cache.go @@ -0,0 +1,152 @@ +package cache + +import ( + "context" + "encoding" + "errors" + "reflect" + "sonet/pkg/plugins/mq" + "sonet/pkg/utils/cache" + "sonet/pkg/utils/logger" + "time" +) + +const ( + flushTopicSuffix = "flush" +) + +type LocalRemoteCache struct { + topic string + flushTopic string + local *cache.KVCache[string, any] + remote Cache + producer mq.Producer + consumer mq.Consumer +} + +type LocalRemoteCacheOptions struct { + Topic string + LocalExpiration time.Duration + CleanupInterval time.Duration + Remote Cache + Producer mq.Producer + Consumer mq.Consumer +} + +func NewLocalRemoteCache(opts LocalRemoteCacheOptions) (*LocalRemoteCache, error) { + expireStore := cache.NewExpireStore[string, any]( + opts.LocalExpiration, + opts.CleanupInterval, + nil, + func(k string) string { + return k + }, + ) + + lrc := &LocalRemoteCache{ + topic: opts.Topic, + flushTopic: opts.Topic + flushTopicSuffix, + local: expireStore, + remote: opts.Remote, + producer: opts.Producer, + consumer: opts.Consumer, + } + return lrc, lrc.init() +} + +func (m *LocalRemoteCache) init() error { + // subscribe mq flush cache msg + return m.consumer.SubscribeBroadcast(m.flushTopic, func(msg *mq.Message) error { + // TODO 判断不删除 local + key := string(msg.Body) + m.local.Delete(key) + logger.Infof("flush cache %s key: %s\n", m.topic, key) + return nil + }) +} + +func (m *LocalRemoteCache) Set(ctx context.Context, key string, value any) error { + err := m.remote.Set(ctx, key, value) + if err != nil { + return err + } + + // set local cache + m.local.SetDefault(key, value) + + // mq flush cache msg + return m.producer.Publish(m.flushTopic, []byte(key)) +} + +func (m *LocalRemoteCache) Load(ctx context.Context, key string, target any) error { + val := m.local.Get(key) + if val != nil { + err := m.copyBinaryMarshal(val, target) + if err != nil { + err2 := m.Del(ctx, key) + if err2 != nil { + return err2 + } + return err + } + return nil + } + logger.Infof("%s load from remote cache: %s\n", m.topic, key) + // load from redis, TODO singleFly + err := m.remote.Load(ctx, key, target) + if err != nil { + return err + } + // cache to local + m.local.SetDefault(key, target) + return nil +} + +func (m *LocalRemoteCache) ForceLoad(ctx context.Context, key string, target any) error { + err := m.remote.Load(ctx, key, target) + if err != nil { + if err == NotExists { + m.local.Delete(key) + } + return err + } + instance := reflect.New(reflect.TypeOf(target).Elem()).Interface() + err = m.copyBinaryMarshal(target, instance) + if err != nil { + return err + } + m.local.SetDefault(key, instance) + return nil +} + +func (m *LocalRemoteCache) copyBinaryMarshal(origin any, target any) error { + marshaler, ok1 := origin.(encoding.BinaryMarshaler) + unmarshaler, ok2 := target.(encoding.BinaryUnmarshaler) + if ok1 && ok2 { + bytes, err := marshaler.MarshalBinary() + if err != nil { + return err + } + if err = unmarshaler.UnmarshalBinary(bytes); err != nil { + return err + } + } else { + return errors.New("value not implement encoding.BinaryMarshaler and encoding.BinaryUnmarshaler") + } + return nil +} + +func (m *LocalRemoteCache) Del(ctx context.Context, keys ...string) error { + err := m.remote.Del(ctx, keys...) + if err != nil { + return err + } + var byteKeys = make([][]byte, 0, len(keys)) + for _, key := range keys { + // klog.Info("del cache: ", key) + m.local.Delete(key) + + byteKeys = append(byteKeys, []byte(key)) + } + return m.producer.MultiPublish(m.flushTopic, byteKeys) +} diff --git a/pkg/plugins/cache/redis.go b/pkg/plugins/cache/redis.go new file mode 100644 index 0000000..5b97e45 --- /dev/null +++ b/pkg/plugins/cache/redis.go @@ -0,0 +1,52 @@ +package cache + +import ( + "context" + "github.com/redis/go-redis/v9" + "time" +) + +type RedisCache struct { + keyPrefix string + rdb *redis.Client +} + +func NewRedisCache(keyPrefix string, rdb *redis.Client) *RedisCache { + return &RedisCache{keyPrefix, rdb} +} + +func (s *RedisCache) Set(ctx context.Context, key string, value any) error { + // redis.writer.go#WriteArg() + return s.rdb.Set(ctx, s.keyPrefix+key, value, 0).Err() +} + +func (s *RedisCache) Load(ctx context.Context, key string, target any) error { + err := s.rdb.Get(ctx, s.keyPrefix+key).Scan(target) + if err == redis.Nil { // key 不存在 + return NotExists + } + if err != nil { + return err + } + return nil +} + +func (s *RedisCache) Del(ctx context.Context, keys ...string) error { + if keys == nil || len(keys) == 0 { + return nil + } + for i, k := range keys { + keys[i] = s.keyPrefix + k + } + return s.rdb.Del(ctx, keys...).Err() +} + +// 重试次数 +var retryTimes = 5 + +// 重试频率 +var retryInterval = time.Millisecond * 50 + +func (s *RedisCache) Lock() { + +} diff --git a/pkg/plugins/mq/mq.go b/pkg/plugins/mq/mq.go new file mode 100644 index 0000000..a1c16a4 --- /dev/null +++ b/pkg/plugins/mq/mq.go @@ -0,0 +1,30 @@ +package mq + +import "time" + +type Message struct { + Id string + Body []byte + Time int64 +} + +type Producer interface { + Publish(topic string, msg []byte) error + + MultiPublish(topic string, msgs [][]byte) error + + DelayPublish(topic string, delay time.Duration, body []byte) error + + Stop() +} + +type Consumer interface { + // Subscribe 订阅 + // channel topic下多个相同channel只收到一次, 不同channel会各收到一次消息副本 + Subscribe(topic string, channel string, handler func(*Message) error) error + + // SubscribeBroadcast 订阅广播消息 + SubscribeBroadcast(topic string, handler func(*Message) error) error + + Stop() +} diff --git a/pkg/plugins/mq/nats.go b/pkg/plugins/mq/nats.go new file mode 100644 index 0000000..e76fc60 --- /dev/null +++ b/pkg/plugins/mq/nats.go @@ -0,0 +1,141 @@ +package mq + +import ( + "github.com/nats-io/nats.go" + "sonet/pkg/utils/logger" + "sync" + "time" +) + +type NatsProducer struct { + options nats.Options + nc *nats.Conn +} + +func NewNatsProducer(options nats.Options) (*NatsProducer, error) { + nc, err := options.Connect() + if err != nil { + return nil, err + } + return &NatsProducer{ + options: options, + nc: nc, + }, nil +} + +func (mq *NatsProducer) Publish(topic string, msg []byte) error { + err := mq.nc.Publish(topic, msg) + if err != nil { + return err + } + return mq.nc.Flush() +} + +func (mq *NatsProducer) MultiPublish(topic string, msgs [][]byte) error { + for _, msg := range msgs { + err := mq.nc.Publish(topic, msg) + if err != nil { + return err + } + } + return mq.nc.Flush() +} + +func (mq *NatsProducer) DelayPublish(topic string, delay time.Duration, body []byte) error { + panic("nats nonsupport delay publish") +} + +func (mq *NatsProducer) Stop() { + err := mq.nc.Flush() + if err != nil { + logger.Error("flush nats error: ", err) + } + mq.nc.Close() +} + +type NatsConsumer struct { + options nats.Options + nc *nats.Conn + consumers map[string]*nats.Subscription + lock *sync.Mutex +} + +func NewNatsConsumer(options nats.Options) (*NatsConsumer, error) { + nc, err := options.Connect() + if err != nil { + return nil, err + } + return &NatsConsumer{ + options: options, + nc: nc, + consumers: make(map[string]*nats.Subscription), + lock: &sync.Mutex{}, + }, nil +} + +func (mq *NatsConsumer) Subscribe(topic string, channel string, handler func(*Message) error) error { + mq.lock.Lock() + defer mq.lock.Unlock() + + if err := mq.unsubscribe(topic); err != nil { + return nil + } + + sub, err := mq.nc.QueueSubscribe(topic, channel, func(msg *nats.Msg) { + message := &Message{Body: msg.Data, Time: time.Now().UnixMilli()} + err := handler(message) + if err != nil { + logger.Error("nats subscribe handle error: ", err) + } + }) + + if err != nil { + return err + } + mq.consumers[topic] = sub + return nil +} + +func (mq *NatsConsumer) SubscribeBroadcast(topic string, handler func(*Message) error) error { + mq.lock.Lock() + defer mq.lock.Unlock() + + if err := mq.unsubscribe(topic); err != nil { + return nil + } + + sub, err := mq.nc.Subscribe(topic, func(msg *nats.Msg) { + message := &Message{Body: msg.Data, Time: time.Now().UnixMilli()} + err := handler(message) + if err != nil { + logger.Error("nats subscribe handle error: ", err) + } + }) + + if err != nil { + return err + } + mq.consumers[topic] = sub + return nil +} + +func (mq *NatsConsumer) Stop() { + mq.lock.Lock() + defer mq.lock.Unlock() + + for topic, sub := range mq.consumers { + err := sub.Unsubscribe() + if err != nil { + logger.Errorf("nats unsubscribe error: topic=%s, err=%v\n", topic, err) + } + } +} + +func (mq *NatsConsumer) unsubscribe(topic string) error { + oldSub := mq.consumers[topic] + if oldSub != nil { + delete(mq.consumers, topic) + return oldSub.Unsubscribe() + } + return nil +} diff --git a/pkg/plugins/mq/nats_jet_stream.go b/pkg/plugins/mq/nats_jet_stream.go new file mode 100644 index 0000000..16fab60 --- /dev/null +++ b/pkg/plugins/mq/nats_jet_stream.go @@ -0,0 +1,190 @@ +package mq + +import ( + "context" + "github.com/google/uuid" + "github.com/nats-io/nats.go" + "github.com/nats-io/nats.go/jetstream" + "regexp" + "sonet/pkg/utils/logger" + "sync" + "time" +) + +var ( + uuidPattern = regexp.MustCompile("^\\w+(-\\w+){4}$") +) + +// NatsJetStreamProducer use for msg ack, msg resend, +type NatsJetStreamProducer struct { + options nats.Options + nc *nats.Conn + js jetstream.JetStream +} + +func NewNatsJetStreamProducer(options nats.Options) *NatsJetStreamProducer { + return &NatsJetStreamProducer{ + options: options, + } +} + +func (mq *NatsJetStreamProducer) Init() error { + nc, err := mq.options.Connect() + if err != nil { + return err + } + mq.nc = nc + + mq.js, err = jetstream.New(mq.nc) + return err +} + +func (mq *NatsJetStreamProducer) Publish(topic string, msg []byte) error { + // TODO 是否ack性能430倍差距,通过配置确定publish方式,大协程池等待ack的网络io(自动重试,失败通知[后续优化]) + //err := mq.nc.Publish(topic, msg) + _, err := mq.js.Publish(context.Background(), topic, msg) + return err +} + +func (mq *NatsJetStreamProducer) MultiPublish(topic string, msgs [][]byte) error { + for _, msg := range msgs { + err := mq.Publish(topic, msg) + if err != nil { + return err + } + } + return nil +} + +func (mq *NatsJetStreamProducer) DelayPublish(topic string, delay time.Duration, body []byte) error { + panic("Nats JetStream nonsupport delay publish") +} + +func (mq *NatsJetStreamProducer) Stop() { + mq.nc.Close() +} + +type NatsJetStreamConsumer struct { + options nats.Options + streamConfig jetstream.StreamConfig + consumerConfig jetstream.ConsumerConfig + + nc *nats.Conn + js jetstream.JetStream + stream jetstream.Stream + + lock *sync.Mutex + consumers map[string]map[string]jetstream.Consumer + consumes map[string]jetstream.ConsumeContext +} + +// NewNatsJetStreamConsumer +// streamConfig name,subjects需配置 +// consumerConfig name, filterSubject 无需配置,在 Subscribe 时配置 +func NewNatsJetStreamConsumer( + options nats.Options, + streamConfig jetstream.StreamConfig, + consumerConfig jetstream.ConsumerConfig) *NatsJetStreamConsumer { + + return &NatsJetStreamConsumer{ + options: options, + streamConfig: streamConfig, + consumerConfig: consumerConfig, + lock: &sync.Mutex{}, + consumers: make(map[string]map[string]jetstream.Consumer), + consumes: make(map[string]jetstream.ConsumeContext), + } +} + +func (mq *NatsJetStreamConsumer) Init(ctx context.Context) error { + nc, err := mq.options.Connect() + if err != nil { + return err + } + mq.nc = nc + mq.js, err = jetstream.New(nc) + + mq.stream, err = mq.js.CreateStream(ctx, mq.streamConfig) + if err != nil { + return err + } + return nil +} + +func (mq *NatsJetStreamConsumer) Subscribe(filterSubject string, channel string, handler func(*Message) error) error { + mq.lock.Lock() + defer mq.lock.Unlock() + + // get consumer + channelConsumers, ok := mq.consumers[channel] + if !ok { + channelConsumers = make(map[string]jetstream.Consumer) + mq.consumers[channel] = channelConsumers + } + consumer, ok := channelConsumers[filterSubject] + if !ok { + consumerConfig := mq.consumerConfig + consumerConfig.Name = channel + consumerConfig.Durable = channel + consumerConfig.FilterSubject = filterSubject + var err error + consumer, err = mq.stream.CreateOrUpdateConsumer(context.Background(), consumerConfig) + if err != nil { + return err + } + channelConsumers[filterSubject] = consumer + } + + // consume message + consumeKey := channel + "/" + filterSubject + consumeContext, ok := mq.consumes[consumeKey] + if ok { + consumeContext.Stop() + delete(mq.consumers, consumeKey) + } + + consumeContext, err := consumer.Consume(func(msg jetstream.Msg) { + message := &Message{ + Body: msg.Data(), + Time: time.Now().UnixMilli(), + } + err := handler(message) + if err == nil { + // ack message + err := msg.Ack() + if err != nil { + logger.Error("ack message error:", err, message) + } + } + }) + if err != nil { + return err + } + mq.consumes[consumeKey] = consumeContext + + return nil +} + +func (mq *NatsJetStreamConsumer) SubscribeBroadcast(filterSubject string, handler func(*Message) error) error { + return mq.Subscribe(filterSubject, uuid.New().String(), handler) +} + +func (mq *NatsJetStreamConsumer) Stop() { + mq.lock.Lock() + defer mq.lock.Unlock() + + for _, consume := range mq.consumes { + consume.Stop() + } + for channel, _ := range mq.consumers { + // 删除 uuid 的广播 consumer + if uuidPattern.MatchString(channel) { + err := mq.js.DeleteConsumer(context.Background(), mq.streamConfig.Name, channel) + if err != nil { + logger.Errorf("delete consumer error, stream=%s, consumer=%s\n", mq.streamConfig.Name, channel) + } + } + } + + mq.nc.Close() +} diff --git a/pkg/protocol/authorize/authorize.go b/pkg/protocol/authorize/authorize.go new file mode 100644 index 0000000..867d25d --- /dev/null +++ b/pkg/protocol/authorize/authorize.go @@ -0,0 +1,36 @@ +package authorize + +import ( + "encoding/base64" + "errors" + "github.com/bytedance/sonic" + "sonet/pkg/utils/security" +) + +type Subject struct { + Uid string `json:"uid,omitempty"` + Username string `json:"username,omitempty"` + Time int64 `json:"time,omitempty"` // ms + Extra map[string]string `json:"extra,omitempty"` +} + +// Verify AuthServer Verify logic +func Verify(aesKey []byte, token string) (subject *Subject, err error) { + bytes, err := base64.URLEncoding.DecodeString(token) + if err != nil { + err = errors.New("token decode fail: " + err.Error()) + return + } + decode, err := security.DecryptAesCBC(bytes, aesKey) + if err != nil { + err = errors.New("invalidate token: " + err.Error()) + return + } + subject = &Subject{} + err = sonic.Unmarshal(decode, subject) + if err != nil { + err = errors.New("token payload decode fail: " + err.Error()) + return + } + return +} diff --git a/pkg/protocol/deliver/deliver.go b/pkg/protocol/deliver/deliver.go new file mode 100644 index 0000000..2acea4b --- /dev/null +++ b/pkg/protocol/deliver/deliver.go @@ -0,0 +1,127 @@ +package deliver + +import ( + "context" + "fmt" + "github.com/golang/protobuf/proto" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "reflect" + "sonet/api/gen/postal" + "sonet/pkg/grpc/discovery" + "sonet/pkg/utils/logger" + "time" +) + +type Status int16 + +const ( + StatusSuccess Status = 1 + StatusError Status = 2 + StatusReceiverOffline Status = 10 +) + +// Deliver n包,通知消息投递 +type Deliver struct { + svcName string + postal postal.PostalClient +} + +func NewDeliver(msgInServiceName string) *Deliver { + return &Deliver{ + svcName: msgInServiceName, + } +} + +func (d *Deliver) InitWithResolver(ctx context.Context, resolver *discovery.Resolver) (err error) { + addr := fmt.Sprintf("%s:///%s", resolver.Scheme(), postal.Postal_ServiceDesc.ServiceName) + conn, err := grpc.DialContext(ctx, addr, + grpc.WithTransportCredentials(insecure.NewCredentials()), + // todo consistent hash lb + // grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingPolicy":"%s"}`, balancer.ConsistentHash)), + ) + if err != nil { + return + } + d.postal = postal.NewPostalClient(conn) + return +} + +func (d *Deliver) InitWithAddr(postalAddr string) (err error) { + // Conn *grpc.ClientConn + conn, err := grpc.Dial(postalAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + return + } + d.postal = postal.NewPostalClient(conn) + return +} + +func (d *Deliver) Deliver(ctx context.Context, msg proto.Message, receiver string, options ...Option) (Status, error) { + return d.deliver0(ctx, msg, []string{receiver}, options...) +} + +func (d *Deliver) DeliverBatch(ctx context.Context, msg proto.Message, receivers []string, options ...Option) (Status, error) { + return d.deliver0(ctx, msg, receivers, options...) +} + +func (d *Deliver) deliver0(ctx context.Context, msg proto.Message, receivers []string, options ...Option) (status Status, err error) { + opts := defaultOptions + if options != nil { + for _, opt := range options { + opt.f(&opts) + } + } + if receivers == nil || len(receivers) == 0 { + // return StatusError, errors.New("receivers is empty") + return StatusSuccess, nil + } + + // encode msg + body, err := proto.Marshal(msg) + if err != nil { + return StatusError, err + } + msgName := reflect.TypeOf(msg).Elem().Name() + message := &postal.Message{ + Time: time.Now().UnixMilli(), + Svc: d.svcName, + Msg: msgName, + Body: body, + } + // deliver to gateway + if len(receivers) == 1 { + // deliver one receiver + reqDeliver := &postal.ReqDeliver{ + Receiver: receivers[0], + Msg: message, + } + res, err := d.postal.Deliver(ctx, reqDeliver) + if err != nil { + return StatusError, err + } + // TODO res code + if res.Ok { + status = StatusSuccess + } else { + status = StatusError + } + logger.Info("deliver result: ", err, res) + } else { + + // deliver batch receiver + req := &postal.ReqDeliverBatch{Receivers: receivers, Msg: message} + res, err := d.postal.DeliverBatch(ctx, req) + if err != nil { + logger.Error("deliver error: ", err) + return StatusError, err + } + if res.Ok { + status = StatusSuccess + } else { + status = StatusError + } + logger.Info("deliver batch result: ", res) + } + return +} diff --git a/pkg/protocol/deliver/options.go b/pkg/protocol/deliver/options.go new file mode 100644 index 0000000..f717cb6 --- /dev/null +++ b/pkg/protocol/deliver/options.go @@ -0,0 +1,19 @@ +package deliver + +var defaultOptions = Options{ + Sync: true, +} + +type Options struct { + Sync bool // default true +} + +type Option struct { + f func(o *Options) +} + +func WithSync(sync bool) Option { + return Option{func(o *Options) { + o.Sync = sync + }} +} diff --git a/pkg/protocol/session/subject.go b/pkg/protocol/session/subject.go new file mode 100644 index 0000000..265735b --- /dev/null +++ b/pkg/protocol/session/subject.go @@ -0,0 +1,49 @@ +package session + +import ( + "context" + "errors" + "google.golang.org/grpc/metadata" +) + +const ( + UidKey = "uid" +) + +var ( + UnauthorizedRequestError = errors.New("unauthorized request") +) + +type RpcSubject struct { + Uid string +} + +func NewRpcSubject(uid string) *RpcSubject { + return &RpcSubject{ + Uid: uid, + } +} + +func PutSubject(ctx context.Context, subject *RpcSubject) context.Context { + return metadata.NewOutgoingContext(ctx, metadata.Pairs(UidKey, subject.Uid)) +} + +func GetSubject(ctx context.Context) (*RpcSubject, error) { + uid, ok := GetUid(ctx) + if !ok { + return nil, UnauthorizedRequestError + } + return &RpcSubject{Uid: uid}, nil +} + +func GetUid(ctx context.Context) (string, bool) { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return "", false + } + uids := md.Get(UidKey) + if len(uids) == 0 { + return "", false + } + return uids[0], true +} diff --git a/pkg/utils/cache/cache.go b/pkg/utils/cache/cache.go new file mode 100644 index 0000000..4cb5bbd --- /dev/null +++ b/pkg/utils/cache/cache.go @@ -0,0 +1,37 @@ +package cache + +import ( + "context" + "encoding" +) + +const NotExists = CacheError("cache: not exists") + +type CacheError string + +func (e CacheError) Error() string { return string(e) } + +type TextSerializable interface { + encoding.TextMarshaler + encoding.TextUnmarshaler +} + +type BinarySerializable interface { + encoding.BinaryMarshaler + encoding.BinaryUnmarshaler +} + +// Cache TODO random Lock(key, timeout), Unlock(key, random) +type Cache interface { + Set(ctx context.Context, key string, value any) error + + Load(ctx context.Context, key string, target any) error + + Del(ctx context.Context, keys ...string) error +} + +type MultiLevelCache interface { + Cache + + ForceLoad(ctx context.Context, key string, target any) error +} diff --git a/pkg/utils/cache/kvcache.go b/pkg/utils/cache/kvcache.go new file mode 100644 index 0000000..4f40aaa --- /dev/null +++ b/pkg/utils/cache/kvcache.go @@ -0,0 +1,59 @@ +package cache + +import ( + "github.com/patrickmn/go-cache" + "time" +) + +const ( + NoExpiration time.Duration = cache.NoExpiration + DefaultExpiration time.Duration = cache.DefaultExpiration +) + +type KVCache[K any, V any] struct { + c *cache.Cache + zeroValue V // 默认值 + k2str func(K) string +} + +func NewKVCache[K any, V any](zeroValue V, k2str func(K) string) *KVCache[K, V] { + return NewExpireStore(cache.NoExpiration, cache.NoExpiration, zeroValue, k2str) +} + +func NewExpireStore[K any, V any](defaultExpiration, cleanupInterval time.Duration, zeroValue V, k2str func(K) string) *KVCache[K, V] { + return &KVCache[K, V]{ + c: cache.New(defaultExpiration, cleanupInterval), + zeroValue: zeroValue, + k2str: k2str, + } +} + +func (s *KVCache[K, V]) SetDefault(k K, v V) { + s.c.SetDefault(s.k2str(k), v) +} + +func (s *KVCache[K, V]) Set(k K, v V, d time.Duration) { + s.c.Set(s.k2str(k), v, d) +} + +func (s *KVCache[K, V]) Get(k K) (zero V) { + v, found := s.c.Get(s.k2str(k)) + if !found { + return s.zeroValue + } + return v.(V) +} + +func (s *KVCache[K, V]) Delete(k K) { + s.c.Delete(s.k2str(k)) +} + +func (s *KVCache[K, V]) ForEach(f func(k string, v V)) { + for k, v := range s.c.Items() { + f(k, v.Object.(V)) + } +} + +func (s *KVCache[K, V]) Count() int { + return s.c.ItemCount() +} diff --git a/pkg/utils/conver/unit_conver.go b/pkg/utils/conver/unit_conver.go index 92b1166..fb50f06 100644 --- a/pkg/utils/conver/unit_conver.go +++ b/pkg/utils/conver/unit_conver.go @@ -3,6 +3,7 @@ package conver import ( "fmt" "github.com/dsnet/golib/unitconv" + "time" ) // ParseDataUnit parse 1Ki -> 1024, 1K -> 1000 @@ -28,3 +29,15 @@ func MustParseDataUnit(unit string) (val float64) { func MustParseDataUnitInt(unit string) (val int) { return int(MustParseDataUnit(unit)) } + +func ParseDuration(s string) (time.Duration, error) { + return time.ParseDuration(s) +} + +func MustParseDuration(s string) time.Duration { + d, err := ParseDuration(s) + if err != nil { + panic(err) + } + return d +} diff --git a/pkg/utils/security/aes.go b/pkg/utils/security/aes.go new file mode 100644 index 0000000..afa3832 --- /dev/null +++ b/pkg/utils/security/aes.go @@ -0,0 +1,116 @@ +package security + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "encoding/base64" + "fmt" + "reflect" +) + +// 参考:https://www.yisu.com/zixun/696240.html +func TestAes() { + defer func() { + if r := recover(); r != nil { + fmt.Println("recover...", r.(error).Error(), reflect.TypeOf(r)) + } + }() + + // aesKeyStr := GetAesKey256() + aesKeyStr := "YYNLVy4qZ+WwkyG7r4kUhMfZu6g9e+" // 0uiba5DAgINl8= + + fmt.Println(aesKeyStr) + key, _ := base64.StdEncoding.DecodeString(aesKeyStr) + fmt.Println(key) + + text := []byte("you are my sunshine!") + encrypt, _ := EncryptAesCBC(text, key) + fmt.Println(encrypt) + fmt.Println(base64.URLEncoding.EncodeToString(text)) + + result, _ := DecryptAesCBC(encrypt, key) + fmt.Println(string(result)) +} + +func GetAesKey256() string { + key := getAesKey(256) + return base64.StdEncoding.EncodeToString(key) +} + +func getAesKey(keySize int) []byte { + var key = make([]byte, keySize/8) + _, err := rand.Read(key) + if err != nil { + panic(err) + } + return key +} + +// EncryptAesCBC +// src -> 要加密的原文 +// key -> 秘钥, 和加密秘钥相同, 大小为: 8byte +func EncryptAesCBC(src, key []byte) ([]byte, error) { + block, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + blockSize := block.BlockSize() + // 对最后一个明文分组进行数据填充 + src = pkcs5Padding(src, blockSize) + // 3.创建一个密码分组为链接模式的,底层使用 DES 加密的 BlockMode 接口 + // 参数 iv 的长度,必须等于 b 的块尺寸 + iv := make([]byte, blockSize, blockSize) + copy(iv, key) + blackMode := cipher.NewCBCEncrypter(block, iv) + // 5.加密连续的数据块 + dst := make([]byte, len(src)) + blackMode.CryptBlocks(dst, src) + return dst, nil +} + +// DecryptAesCBC +// src -> 要解密的密文 +// key -> 秘钥, 和加密秘钥相同, 大小为: 8byte +func DecryptAesCBC(src, key []byte) ([]byte, error) { + // 1. 创建并返回一个使用DES算法的cipher.Block接口 + block, err := aes.NewCipher(key) + // 2. 判断是否创建成功 + if err != nil { + return nil, err + } + blockSize := block.BlockSize() + // 3. 创建一个密码分组为链接模式的, 底层使用DES解密的BlockMode接口 + iv := make([]byte, blockSize, blockSize) + copy(iv, key) + blockMode := cipher.NewCBCDecrypter(block, iv) + // 4. 解密数据 + dst := src + blockMode.CryptBlocks(src, dst) + // 5. 去掉最后一组填充的数据 + dst = pkcs5UnPadding(dst) + // 6. 返回结果 + return dst, nil +} + +// PKCS5Padding 使用pks5的方式填充 +func pkcs5Padding(ciphertext []byte, blockSize int) []byte { + // 1. 计算最后一个分组缺多少个字节 + padding := blockSize - (len(ciphertext) % blockSize) + // 2. 创建一个大小为padding的切片, 每个字节的值为padding + padText := bytes.Repeat([]byte{byte(padding)}, padding) + // 3. 将padText添加到原始数据的后边, 将最后一个分组缺少的字节数补齐 + newText := append(ciphertext, padText...) + return newText +} + +// PKCS5UnPadding 删除pks5填充的尾部数据 +func pkcs5UnPadding(origData []byte) []byte { + // 1. 计算数据的总长度 + length := len(origData) + // 2. 根据填充的字节值得到填充的次数 + number := int(origData[length-1]) + // 3. 将尾部填充的number个字节去掉 + return origData[:(length - number)] +} diff --git a/pkg/utils/shutdown/signal.go b/pkg/utils/shutdown/signal.go index a7ff8ba..b8a0630 100644 --- a/pkg/utils/shutdown/signal.go +++ b/pkg/utils/shutdown/signal.go @@ -4,6 +4,7 @@ import ( "log" "os" "os/signal" + "sonet/pkg/utils/logger" "sync" "sync/atomic" "syscall" @@ -26,7 +27,7 @@ func Await() { signal.Notify(sigChan, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) s := <-sigChan // 监听到关闭信号 - log.Println("catch exit signal: ", s) + logger.Info("catch exit signal: ", s) var success, fail int32 for _, hook := range shutdownHooks { @@ -44,7 +45,7 @@ func Await() { }() } - log.Printf("execute shutdown hook %d success, %d failed\n", success, fail) + logger.Info("execute %d shutdown hook %d ok, %d failed\n", success+fail, success, fail) } //func Shutdown() { diff --git a/pkg/utils/strs/random.go b/pkg/utils/strs/random.go new file mode 100644 index 0000000..964f8a9 --- /dev/null +++ b/pkg/utils/strs/random.go @@ -0,0 +1,85 @@ +// reference go-zero/core/stringx/random.go + +package strs + +import ( + crand "crypto/rand" + "fmt" + "math/rand" + "sync" + "time" +) + +const ( + letterBytes = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" + letterIdxBits = 6 // 6 bits to represent a letter index + idLen = 8 + defaultRandLen = 8 + letterIdxMask = 1<= 0; { + if remain == 0 { + cache, remain = src.Int63(), letterIdxMax + } + if idx := int(cache & letterIdxMask); idx < len(letterBytes) { + b[i] = letterBytes[idx] + i-- + } + cache >>= letterIdxBits + remain-- + } + + return string(b) +} + +// Seed sets the seed to seed. +func Seed(seed int64) { + src.Seed(seed) +} diff --git a/pkg/utils/strs/variable.go b/pkg/utils/strs/variable.go new file mode 100644 index 0000000..7d15490 --- /dev/null +++ b/pkg/utils/strs/variable.go @@ -0,0 +1,74 @@ +package strs + +import "strings" + +// SnakeString 驼峰转蛇形 +func SnakeString(s string) string { + data := make([]byte, 0, len(s)*2) + j := false + num := len(s) + for i := 0; i < num; i++ { + d := s[i] + // or通过ASCII码进行大小写的转化 + // 65-90(A-Z),97-122(a-z) + //判断如果字母为大写的A-Z就在前面拼接一个_ + if i > 0 && d >= 'A' && d <= 'Z' && j { + data = append(data, '_') + } + if d != '_' { + j = true + } + data = append(data, d) + } + //ToLower把大写字母统一转小写 + return strings.ToLower(string(data[:])) +} + +// CamelString 蛇形转驼峰 +func CamelString(s string) string { + data := make([]byte, 0, len(s)) + j := false + k := false + num := len(s) - 1 + for i := 0; i <= num; i++ { + d := s[i] + if k == false && d >= 'A' && d <= 'Z' { + k = true + } + if d >= 'a' && d <= 'z' && (j || k == false) { + d = d - 32 + j = false + k = true + } + if k && d == '_' && num > i && s[i+1] >= 'a' && s[i+1] <= 'z' { + j = true + continue + } + data = append(data, d) + } + return string(data[:]) +} + +// LowerInitialLetter 首字母小写 +func LowerInitialLetter(s string) string { + if s == "" { + return s + } + letters := []rune(s) + if letters[0] >= 'A' && letters[0] <= 'Z' { + letters[0] += 32 + } + return string(letters) +} + +// UpperInitialLetter 首字母大写 +func UpperInitialLetter(s string) string { + if s == "" { + return s + } + letters := []rune(s) + if letters[0] >= 'a' && letters[0] <= 'z' { + letters[0] -= 32 + } + return string(letters) +}