47 changed files with 2357 additions and 117 deletions
@ -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 |
||||
@ -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() |
||||
} |
||||
@ -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 |
||||
@ -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 |
||||
} |
||||
@ -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())
|
||||
|
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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" |
||||
} |
||||
@ -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,15 +0,0 @@
|
||||
package session |
||||
|
||||
import ( |
||||
"sync" |
||||
) |
||||
|
||||
// NetAccount 已认证的长连接用户
|
||||
type NetAccount struct { |
||||
Id int64 |
||||
UserName string |
||||
Avatar string |
||||
Client *NetClient |
||||
|
||||
Lock *sync.Mutex |
||||
} |
||||
@ -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, |
||||
} |
||||
} |
||||
@ -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)) |
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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 { |
||||
} |
||||
@ -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) |
||||
} |
||||
@ -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() { |
||||
|
||||
} |
||||
@ -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() |
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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() |
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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 |
||||
}} |
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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 |
||||
} |
||||
@ -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() |
||||
} |
||||
@ -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)] |
||||
} |
||||
@ -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) |
||||
} |
||||
@ -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…
Reference in new issue