Browse Source

deliver msg

master
tangmingyou 3 years ago
parent
commit
00f70b14d0
  1. 22
      api/postal.proto
  2. 1
      cmd/auth/config.toml
  3. 14
      cmd/auth/main.go
  4. 33
      cmd/chat/config.toml
  5. 46
      cmd/chat/main.go
  6. 24
      cmd/gateway_http/config.toml
  7. 113
      cmd/gateway_http/main.go
  8. 3
      cmd/gateway_ws/config.toml
  9. 66
      cmd/gateway_ws/main.go
  10. 30
      cmd/mahjong/main.go
  11. 12
      go.mod
  12. 24
      go.sum
  13. 66
      internal/auth/data/auth_user_dao.go
  14. 18
      internal/auth/data/t_user.go
  15. 109
      internal/auth/logic/auth_server.go
  16. 69
      internal/chat/logic/chat_server.go
  17. 55
      internal/gateway_http/config/auth_filter.go
  18. 1
      internal/gateway_http/logic/http_server.go
  19. 94
      internal/gateway_ws/server/ws_server.go
  20. 15
      internal/gateway_ws/session/net_account.go
  21. 27
      internal/gateway_ws/session/net_client.go
  22. 15
      internal/gateway_ws/session/net_subject.go
  23. 77
      internal/postal/logic/postal_cluster_server.go
  24. 194
      internal/postal/logic/postal_server.go
  25. 1
      pkg/config/grpc_options.go
  26. 1
      pkg/deliver/deliver.go
  27. 41
      pkg/grpc/client/direct_client_factory.go
  28. 27
      pkg/grpc/discovery/register.go
  29. 34
      pkg/grpc/generic/generic_client.go
  30. 11
      pkg/grpc/generic/generic_client_factory.go
  31. 46
      pkg/plugins/cache/cache.go
  32. 152
      pkg/plugins/cache/multi_cache.go
  33. 52
      pkg/plugins/cache/redis.go
  34. 30
      pkg/plugins/mq/mq.go
  35. 141
      pkg/plugins/mq/nats.go
  36. 190
      pkg/plugins/mq/nats_jet_stream.go
  37. 36
      pkg/protocol/authorize/authorize.go
  38. 127
      pkg/protocol/deliver/deliver.go
  39. 19
      pkg/protocol/deliver/options.go
  40. 49
      pkg/protocol/session/subject.go
  41. 37
      pkg/utils/cache/cache.go
  42. 59
      pkg/utils/cache/kvcache.go
  43. 13
      pkg/utils/conver/unit_conver.go
  44. 116
      pkg/utils/security/aes.go
  45. 5
      pkg/utils/shutdown/signal.go
  46. 85
      pkg/utils/strs/random.go
  47. 74
      pkg/utils/strs/variable.go

22
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;
}

1
cmd/auth/config.toml

@ -1,4 +1,5 @@
[app]
aesTokenKey = "9Nz3Y6DES3msAFndz4QsJAECUIFkf+KKaRa+jRnNALk="
[grpc]
address = ":7020"

14
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 {

33
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

46
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()
}

24
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

113
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
}

3
cmd/gateway_ws/config.toml

@ -1,5 +1,8 @@
[app]
httpPort = 7001
subjectCacheTopic = "wsgate:subject:"
subjectLrcExpiration = "10m"
subjectLrcCleanupInterval = "5m"
[grpc]
address = ":7010"

66
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
}

30
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())
}

12
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

24
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=

66
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
}

18
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"
}

109
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
}

69
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) {

55
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)
}

1
internal/gateway_http/logic/http_server.go

@ -0,0 +1 @@
package logic

94
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 // <string, *session.NetSubject> 当前连接用户,内存缓存
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) {

15
internal/gateway_ws/session/net_account.go

@ -1,15 +0,0 @@
package session
import (
"sync"
)
// NetAccount 已认证的长连接用户
type NetAccount struct {
Id int64
UserName string
Avatar string
Client *NetClient
Lock *sync.Mutex
}

27
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())

15
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,
}
}

77
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))
}

194
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
}

1
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
}

1
pkg/deliver/deliver.go

@ -1 +0,0 @@
package deliver

41
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
}

27
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
}

34
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
}

11
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...)

46
pkg/plugins/cache/cache.go vendored

@ -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 {
}

152
pkg/plugins/cache/multi_cache.go vendored

@ -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)
}

52
pkg/plugins/cache/redis.go vendored

@ -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() {
}

30
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()
}

141
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
}

190
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()
}

36
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
}

127
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
}

19
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
}}
}

49
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
}

37
pkg/utils/cache/cache.go vendored

@ -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
}

59
pkg/utils/cache/kvcache.go vendored

@ -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()
}

13
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
}

116
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)]
}

5
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() {

85
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<<letterIdxBits - 1 // All 1-bits, as many as letterIdxBits
letterIdxMax = 63 / letterIdxBits // # of letter indices fitting in 63 bits
)
var src = newLockedSource(time.Now().UnixNano())
type lockedSource struct {
source rand.Source
lock sync.Mutex
}
func newLockedSource(seed int64) *lockedSource {
return &lockedSource{
source: rand.NewSource(seed),
}
}
func (ls *lockedSource) Int63() int64 {
ls.lock.Lock()
defer ls.lock.Unlock()
return ls.source.Int63()
}
func (ls *lockedSource) Seed(seed int64) {
ls.lock.Lock()
defer ls.lock.Unlock()
ls.source.Seed(seed)
}
// Rand returns a random string.
func Rand() string {
return Randn(defaultRandLen)
}
// RandId returns a random id string.
func RandId() string {
b := make([]byte, idLen)
_, err := crand.Read(b)
if err != nil {
return Randn(idLen)
}
return fmt.Sprintf("%x%x%x%x", b[0:2], b[2:4], b[4:6], b[6:8])
}
// Randn returns a random string with length n.
func Randn(n int) string {
b := make([]byte, n)
// A src.Int63() generates 63 random bits, enough for letterIdxMax characters!
for i, cache, remain := n-1, src.Int63(), letterIdxMax; i >= 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)
}

74
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)
}
Loading…
Cancel
Save