20 changed files with 820 additions and 50 deletions
@ -1 +1,12 @@
|
||||
package main |
||||
|
||||
import ( |
||||
"encoding/binary" |
||||
"fmt" |
||||
) |
||||
|
||||
func main() { |
||||
buf := make([]byte, 4) |
||||
binary.BigEndian.PutUint32(buf, 12) |
||||
fmt.Printf("%v\n", buf) |
||||
} |
||||
|
||||
@ -1,7 +1,25 @@
|
||||
package sonet |
||||
package main |
||||
|
||||
import "os" |
||||
|
||||
//go:generate go run generate.go
|
||||
//go:generate protoc --go_out=./api/gen --go-grpc_out=./api/gen ./api/*.proto
|
||||
//go:generate pbjs -t static-module -w es6 -o api/genjs/postal.js api/postal.proto --no-service
|
||||
//go:generate pbjs -t static-module -w es6 -o api/genjs/auth.js api/auth.proto --no-service
|
||||
//go:generate pbjs -t static-module -w es6 -o api/genjs/chat.js api/chat.proto --no-service
|
||||
//go:generate pbjs -t static-module -w es6 -o api/genjs/mahjong.js api/mahjong.proto --no-service
|
||||
|
||||
func main() { |
||||
beforeGenerate() |
||||
} |
||||
|
||||
func beforeGenerate() { |
||||
err := os.MkdirAll("api/gen", os.ModePerm) |
||||
if err != nil { |
||||
panic(err) |
||||
} |
||||
err = os.MkdirAll("api/genjs", os.ModePerm) |
||||
if err != nil { |
||||
panic(err) |
||||
} |
||||
} |
||||
|
||||
@ -0,0 +1,63 @@
|
||||
package server |
||||
|
||||
import ( |
||||
"fmt" |
||||
"github.com/gin-gonic/gin" |
||||
"github.com/gorilla/websocket" |
||||
"net/http" |
||||
"sonet/pkg/utils/conver" |
||||
"sonet/pkg/utils/resp" |
||||
"time" |
||||
) |
||||
|
||||
type HttpServer struct { |
||||
connHandler *ConnHandler |
||||
} |
||||
|
||||
func NewHttpServer(connHandler *ConnHandler) *HttpServer { |
||||
return &HttpServer{ |
||||
connHandler: connHandler, |
||||
} |
||||
} |
||||
|
||||
func (s *HttpServer) Run(port int) error { |
||||
server := gin.Default() |
||||
server.GET("/ws", s.upgrade) |
||||
return server.Run(fmt.Sprintf(":%d", port)) |
||||
} |
||||
|
||||
var ( |
||||
HandshakeTimeout = 3 * time.Second |
||||
ReadDeadline = 5 * time.Second |
||||
WriteDeadline = 5 * time.Second |
||||
// PongWait Time allowed to read the next pong message from the peer.
|
||||
PongWait = 60 * time.Second |
||||
|
||||
MaxMessageSize = conver.MustParseDataUnitInt("4M") |
||||
ReadBufferSize = conver.MustParseDataUnitInt("4Ki") |
||||
WriteBufferSize = conver.MustParseDataUnitInt("4Ki") |
||||
) |
||||
|
||||
func (s *HttpServer) upgrade(ctx *gin.Context) { |
||||
upgrader := websocket.Upgrader{ |
||||
ReadBufferSize: ReadBufferSize, |
||||
WriteBufferSize: WriteBufferSize, |
||||
HandshakeTimeout: HandshakeTimeout, |
||||
CheckOrigin: func(r *http.Request) bool { |
||||
return true |
||||
}, |
||||
} |
||||
conn, err := upgrader.Upgrade(ctx.Writer, ctx.Request, nil) |
||||
if err != nil { |
||||
ctx.JSON(http.StatusOK, resp.Error(err.Error())) |
||||
return |
||||
} |
||||
|
||||
// https://github.com/gorilla/websocket/blob/a68708917c6a4f06314ab4e52493cc61359c9d42/examples/chat/conn.go#L50
|
||||
conn.SetReadLimit(int64(MaxMessageSize)) |
||||
conn.SetPongHandler(func(string) error { |
||||
return conn.SetReadDeadline(time.Now().Add(PongWait)) |
||||
}) |
||||
|
||||
go s.connHandler.handleConn(conn) |
||||
} |
||||
@ -0,0 +1,89 @@
|
||||
package server |
||||
|
||||
import ( |
||||
"context" |
||||
"fmt" |
||||
"github.com/gorilla/websocket" |
||||
"runtime/debug" |
||||
"sonet/internal/gateway_ws/session" |
||||
"sonet/pkg/grpc/generic" |
||||
"sonet/pkg/protocol" |
||||
"sonet/pkg/utils/logger" |
||||
) |
||||
|
||||
type ConnHandler struct { |
||||
grpcFactory *generic.GrpcGenericClientFactory |
||||
} |
||||
|
||||
func NewConnHandler(grpcFactory *generic.GrpcGenericClientFactory) *ConnHandler { |
||||
return &ConnHandler{ |
||||
grpcFactory: grpcFactory, |
||||
} |
||||
} |
||||
|
||||
func (c *ConnHandler) handleConn(conn *websocket.Conn) { |
||||
client := session.NewNetClient(conn, ReadDeadline, WriteDeadline) |
||||
|
||||
defer func() { |
||||
// 捕获其他错误
|
||||
if r := recover(); r != nil { |
||||
logger.Error("NetClient recover error: ", r) |
||||
// 输出堆栈信息
|
||||
logger.Error("NetClient recover error stack: ", string(debug.Stack())) |
||||
} |
||||
}() |
||||
|
||||
defer client.Close() |
||||
|
||||
for { |
||||
message, err := client.ReadMessage() |
||||
if err != nil { |
||||
// TODO 连接关闭,mq发送关闭事件
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway) { |
||||
logger.Error("unexpected close error: ", err) |
||||
} else { |
||||
// 读失败
|
||||
logger.Errorf("NetClient ReadMessage error: %T, %v", err, err) |
||||
} |
||||
return |
||||
} |
||||
|
||||
payload, err := protocol.Decode(message) |
||||
if err != nil { |
||||
logger.Errorf("decode message error: len=%d", len(message), err) |
||||
return |
||||
} |
||||
fmt.Printf("%v\n", payload) |
||||
|
||||
// grpc generic call
|
||||
ctx := context.Background() |
||||
grpcClient, err := c.grpcFactory.GetClient(ctx, payload.Header.Svc) |
||||
if err != nil { |
||||
logger.Error("get grpc generic client error: ", err) |
||||
continue |
||||
} |
||||
resp, err := grpcClient.InvokeUnary(ctx, payload.Header.Target, payload.Body) |
||||
if err != nil { |
||||
logger.Error("grpc generic call error: ", err) |
||||
continue |
||||
} |
||||
fmt.Println(resp) |
||||
|
||||
// write response
|
||||
payload.Header.Type = protocol.TypeResponse |
||||
payload.Header.Svc = "" |
||||
payload.Header.Target = "" |
||||
payload.Body, err = resp.Marshal() |
||||
if err != nil { |
||||
logger.Error("generic call response marshal error: ", err) |
||||
continue |
||||
} |
||||
resMessage, err := protocol.Encode(payload) |
||||
if err != nil { |
||||
logger.Error("grpc generic call error: ", err) |
||||
continue |
||||
} |
||||
client.MustWrite(resMessage) |
||||
} |
||||
|
||||
} |
||||
@ -0,0 +1,15 @@
|
||||
package session |
||||
|
||||
import ( |
||||
"sync" |
||||
) |
||||
|
||||
// NetAccount 已认证的长连接用户
|
||||
type NetAccount struct { |
||||
Id int64 |
||||
UserName string |
||||
Avatar string |
||||
Client *NetClient |
||||
|
||||
Lock *sync.Mutex |
||||
} |
||||
@ -0,0 +1,77 @@
|
||||
package session |
||||
|
||||
import ( |
||||
"fmt" |
||||
"github.com/gorilla/websocket" |
||||
"sonet/pkg/utils/logger" |
||||
"sync" |
||||
"sync/atomic" |
||||
"time" |
||||
) |
||||
|
||||
// NetClient 长连接客户端
|
||||
type NetClient struct { |
||||
Conn *websocket.Conn |
||||
writeDeadline, readDeadline time.Duration |
||||
Online *atomic.Bool |
||||
Account *NetAccount |
||||
WriteLock *sync.Mutex |
||||
} |
||||
|
||||
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, |
||||
} |
||||
} |
||||
|
||||
func (c *NetClient) Close() { |
||||
err := c.Conn.Close() |
||||
if err != nil { |
||||
logger.Error("NetClient Close error: ", err) |
||||
} |
||||
} |
||||
|
||||
func (c *NetClient) Write(bytes []byte) (err error) { |
||||
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()) |
||||
return |
||||
} |
||||
err = c.Conn.WriteMessage(websocket.BinaryMessage, bytes) |
||||
if err != nil { |
||||
err = fmt.Errorf("NetClient write message error: %s", err.Error()) |
||||
} |
||||
return |
||||
} |
||||
|
||||
func (c *NetClient) MustWrite(bytes []byte) { |
||||
err := c.Write(bytes) |
||||
if err != nil { |
||||
logger.Error(err) |
||||
return |
||||
} |
||||
} |
||||
|
||||
func (c *NetClient) ReadMessage() (bytes []byte, err error) { |
||||
err = c.Conn.SetReadDeadline(time.Now().Add(c.readDeadline)) |
||||
if err != nil { |
||||
return |
||||
} |
||||
|
||||
var messageType int |
||||
messageType, bytes, err = c.Conn.ReadMessage() |
||||
if messageType != websocket.BinaryMessage { |
||||
err = fmt.Errorf("only support websocket binary message") |
||||
return |
||||
} |
||||
return |
||||
} |
||||
@ -0,0 +1,23 @@
|
||||
package session |
||||
|
||||
import "github.com/bytedance/sonic" |
||||
|
||||
// Subject 消息传递主体
|
||||
type Subject struct { |
||||
Uid string `json:"uid,omitempty"` |
||||
Online int8 `json:"online,omitempty"` // 0离线,1在线
|
||||
Time int64 `json:"time,omitempty"` |
||||
Gate string `json:"gate,omitempty"` // 连接的网关addr
|
||||
} |
||||
|
||||
func (sub *Subject) MarshalBinary() (data []byte, err error) { |
||||
bytes, err := sonic.Marshal(sub) |
||||
if err != nil { |
||||
return nil, err |
||||
} |
||||
return bytes, nil |
||||
} |
||||
|
||||
func (sub *Subject) UnmarshalBinary(data []byte) error { |
||||
return sonic.Unmarshal(data, sub) |
||||
} |
||||
@ -1,4 +0,0 @@
|
||||
package ws |
||||
|
||||
type WebsocketServer struct { |
||||
} |
||||
@ -0,0 +1,145 @@
|
||||
package generic |
||||
|
||||
import ( |
||||
"context" |
||||
"fmt" |
||||
"github.com/jhump/protoreflect/desc" |
||||
"github.com/jhump/protoreflect/dynamic" |
||||
"github.com/jhump/protoreflect/dynamic/grpcdynamic" |
||||
"github.com/jhump/protoreflect/grpcreflect" |
||||
"google.golang.org/grpc" |
||||
"google.golang.org/grpc/reflection/grpc_reflection_v1alpha" |
||||
"sonet/pkg/grpc/generic/desc_source" |
||||
"sonet/pkg/utils/logger" |
||||
"sync" |
||||
) |
||||
|
||||
type GrpcGenericClient struct { |
||||
serviceName string |
||||
conn *grpc.ClientConn |
||||
|
||||
descSource desc_source.DescriptorSource |
||||
serviceDesc *desc.ServiceDescriptor |
||||
callerCache *sync.Map |
||||
} |
||||
|
||||
func NewGpcGenericClient(serviceName string, conn *grpc.ClientConn) *GrpcGenericClient { |
||||
return &GrpcGenericClient{ |
||||
serviceName: serviceName, |
||||
conn: conn, |
||||
} |
||||
} |
||||
|
||||
func (c *GrpcGenericClient) Init(ctx context.Context) (err error) { |
||||
// fetch service description
|
||||
refClient := grpcreflect.NewClientV1Alpha(ctx, grpc_reflection_v1alpha.NewServerReflectionClient(c.conn)) |
||||
desc_source.DescriptorSourceFromServer(ctx, refClient) |
||||
dsc, e1 := c.descSource.FindSymbol(c.serviceName) |
||||
if e1 != nil { |
||||
err = fmt.Errorf("service %s not found", c.serviceName) |
||||
if desc_source.IsNotFoundError(e1) { |
||||
logger.Errorf("target server not expose service %q in FindSymbol", c.serviceName) |
||||
return |
||||
} |
||||
logger.Errorf("failed to query for service descriptor %q: %v", c.serviceName, e1) |
||||
return |
||||
} |
||||
|
||||
sd, ok := dsc.(*desc.ServiceDescriptor) |
||||
if !ok { |
||||
err = fmt.Errorf("service %s not found", c.serviceName) |
||||
logger.Errorf("target server not expose service %q", c.serviceName) |
||||
return |
||||
} |
||||
c.serviceDesc = sd |
||||
c.callerCache = &sync.Map{} |
||||
return |
||||
} |
||||
|
||||
func (c *GrpcGenericClient) ServiceName() string { |
||||
return c.serviceName |
||||
} |
||||
|
||||
type methodCaller struct { |
||||
Mtd *desc.MethodDescriptor |
||||
MsgFactory *dynamic.MessageFactory |
||||
Stub grpcdynamic.Stub |
||||
} |
||||
|
||||
func (c *GrpcGenericClient) InvokeUnary(ctx context.Context, method string, reqBytes []byte, 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()) |
||||
if err = reqMessage.(*dynamic.Message).Unmarshal(reqBytes); err != nil { |
||||
err = fmt.Errorf("unmarshal req bytes error: %s", err.Error()) |
||||
return |
||||
} |
||||
|
||||
res, err := caller.Stub.InvokeRpc(ctx, caller.Mtd, reqMessage, opts...) |
||||
if err != nil { |
||||
return |
||||
} |
||||
resp = res.(*dynamic.Message) |
||||
return |
||||
} |
||||
|
||||
// getMethodCaller load generic resource from cache
|
||||
func (c *GrpcGenericClient) getMethodCaller(method string) (caller *methodCaller, err error) { |
||||
val, ok := c.callerCache.Load(method) |
||||
if ok { |
||||
caller = val.(*methodCaller) |
||||
} else { |
||||
// method desc
|
||||
caller = &methodCaller{} |
||||
caller.Mtd = c.serviceDesc.FindMethodByName(method) |
||||
if caller.Mtd == nil { |
||||
logger.Errorf("service %q does not include a method named %q", c.serviceName, method) |
||||
err = fmt.Errorf("method %s not found", method) |
||||
return |
||||
} |
||||
|
||||
// message factory
|
||||
var ext dynamic.ExtensionRegistry |
||||
if err = c.fetchAllExtensions(&ext, caller.Mtd.GetInputType()); err != nil { |
||||
return |
||||
} |
||||
if err = c.fetchAllExtensions(&ext, caller.Mtd.GetOutputType()); err != nil { |
||||
return |
||||
} |
||||
caller.MsgFactory = dynamic.NewMessageFactoryWithExtensionRegistry(&ext) |
||||
|
||||
// stub
|
||||
caller.Stub = grpcdynamic.NewStubWithMessageFactory(c.conn, caller.MsgFactory) |
||||
c.callerCache.Store(method, caller) |
||||
} |
||||
return |
||||
} |
||||
|
||||
func (c *GrpcGenericClient) fetchAllExtensions(ext *dynamic.ExtensionRegistry, md *desc.MessageDescriptor) (err error) { |
||||
msgTypeName := md.GetFullyQualifiedName() |
||||
if len(md.GetExtensionRanges()) > 0 { |
||||
fds, err := c.descSource.AllExtensionsForType(msgTypeName) |
||||
if err != nil { |
||||
return fmt.Errorf("failed to query for extensions of type %s: %v", msgTypeName, err) |
||||
} |
||||
for _, fd := range fds { |
||||
if err := ext.AddExtension(fd); err != nil { |
||||
return fmt.Errorf("could not register extension %s of type %s: %v", fd.GetFullyQualifiedName(), msgTypeName, err) |
||||
} |
||||
} |
||||
} |
||||
// recursively fetch extensions for the types of any message fields
|
||||
for _, fd := range md.GetFields() { |
||||
if fd.GetMessageType() != nil { |
||||
err := c.fetchAllExtensions(ext, fd.GetMessageType()) |
||||
if err != nil { |
||||
return err |
||||
} |
||||
} |
||||
} |
||||
return nil |
||||
} |
||||
@ -0,0 +1,51 @@
|
||||
package generic |
||||
|
||||
import ( |
||||
"context" |
||||
"fmt" |
||||
"google.golang.org/grpc" |
||||
"sonet/pkg/grpc/discovery" |
||||
"sync" |
||||
) |
||||
|
||||
type GrpcGenericClientFactory struct { |
||||
resolver *discovery.Resolver |
||||
defaultOpts []grpc.DialOption |
||||
clientCache *sync.Map |
||||
} |
||||
|
||||
func NewGpcGenericClientFactory(resolver *discovery.Resolver, defaultOpts ...grpc.DialOption) *GrpcGenericClientFactory { |
||||
return &GrpcGenericClientFactory{ |
||||
resolver: resolver, |
||||
defaultOpts: defaultOpts, |
||||
} |
||||
} |
||||
|
||||
func (f *GrpcGenericClientFactory) Init() { |
||||
f.clientCache = &sync.Map{} |
||||
} |
||||
|
||||
func (f *GrpcGenericClientFactory) NewClient(ctx context.Context, serviceName string, opts ...grpc.DialOption) (client *GrpcGenericClient, err error) { |
||||
addr := fmt.Sprintf("%s:///%s", f.resolver.Scheme(), serviceName) |
||||
dialOpts := make([]grpc.DialOption, len(f.defaultOpts)+len(opts)) |
||||
dialOpts = append(dialOpts, f.defaultOpts...) |
||||
dialOpts = append(dialOpts, opts...) |
||||
conn, err := grpc.DialContext(ctx, addr, dialOpts...) |
||||
client = NewGpcGenericClient(serviceName, conn) |
||||
err = client.Init(ctx) |
||||
return |
||||
} |
||||
|
||||
func (f *GrpcGenericClientFactory) GetClient(ctx context.Context, serviceName string, opts ...grpc.DialOption) (client *GrpcGenericClient, err error) { |
||||
val, ok := f.clientCache.Load(serviceName) |
||||
if ok { |
||||
client = val.(*GrpcGenericClient) |
||||
return |
||||
} |
||||
client, err = f.NewClient(ctx, serviceName, opts...) |
||||
if err != nil { |
||||
return |
||||
} |
||||
f.clientCache.Store(serviceName, client) |
||||
return |
||||
} |
||||
@ -1,38 +1,123 @@
|
||||
package protocol |
||||
|
||||
import ( |
||||
"encoding/binary" |
||||
"fmt" |
||||
) |
||||
|
||||
// protocol
|
||||
// 1byte: magic: 99
|
||||
// 1byte: 1 request, 2 response, 3 event, 4 error
|
||||
// 1byte: type 1 request, 2 response, 3 event, 4 error
|
||||
// 1byte: status: 20-OK, 30-CLIENT_TIMEOUT, 31-SERVER_TIMEOUT, 40-BAD_REQUEST, 41-BAD_RESPONSE, 44-SERVICE_NOT_FOUND, 50-CLIENT_ERROR, 51-SERVER_ERROR, 52-SERVICE_ERROR
|
||||
// 4bit: svc,method url desc: 1 name, 2 number
|
||||
// 4bit: serialize: 1proto, 2json
|
||||
// 4bit: urlType svc,method url desc: 1 name, 2 number
|
||||
// 4bit: serializeType: 1proto, 2json
|
||||
// 4byte: seqId
|
||||
// string: service, method / 4byte svc, 4byte method/4byte notice
|
||||
// proto bytes / json bytes
|
||||
|
||||
const ( |
||||
Magic int8 = 99 |
||||
Magic byte = 99 |
||||
|
||||
TypeRequest int8 = 1 |
||||
TypeResponse int8 = 2 |
||||
TypeNotice int8 = 3 |
||||
TypeError int8 = 4 |
||||
TypeRequest byte = 1 |
||||
TypeResponse byte = 2 |
||||
TypeNotice byte = 3 |
||||
TypeError byte = 4 |
||||
) |
||||
|
||||
type Header struct { |
||||
Magic int8 |
||||
Type int8 // 1 request, 2 response, 3 event, 4 error
|
||||
Status int8 |
||||
UrlType int8 |
||||
SerializeType int8 |
||||
Magic byte |
||||
Type byte // 1 request, 2 response, 3 event, 4 error
|
||||
Status byte |
||||
UrlType byte // 4bit: svc,method url desc: 1 name, 2 number
|
||||
SerializeType byte // 4bit: body serialize: 1proto, 2json
|
||||
SeqId int32 |
||||
SvcNo int32 |
||||
MethodNo int32 |
||||
Svc string |
||||
Method string |
||||
SvcNo int32 // optional 1
|
||||
TargetNo int32 // optional 2
|
||||
Svc string // optional 1
|
||||
Target string // optional 2 method/message
|
||||
} |
||||
|
||||
type Payload struct { |
||||
Header Header |
||||
Body []byte |
||||
} |
||||
|
||||
func Decode(bytes []byte) (payload *Payload, err error) { |
||||
header := Header{} |
||||
|
||||
header.Magic = bytes[0] |
||||
if header.Magic != Magic { |
||||
err = fmt.Errorf("unknown magic: %d", header.Magic) |
||||
return |
||||
} |
||||
header.Type = bytes[1] |
||||
header.Status = bytes[2] |
||||
header.UrlType = bytes[3] >> 4 |
||||
header.SerializeType = bytes[3] & 0xF |
||||
header.SeqId = int32(binary.BigEndian.Uint32(bytes[4:8])) |
||||
|
||||
var cursor int |
||||
switch header.UrlType { |
||||
case 1: |
||||
svcLen := int(binary.BigEndian.Uint32(bytes[8:12])) |
||||
header.Svc = string(bytes[12 : 12+svcLen]) |
||||
cursor = 12 + svcLen |
||||
|
||||
targetLen := int(binary.BigEndian.Uint32(bytes[cursor : cursor+4])) |
||||
header.Target = string(bytes[cursor+4 : cursor+4+targetLen]) |
||||
cursor = cursor + 4 + targetLen |
||||
case 2: |
||||
header.SvcNo = int32(binary.BigEndian.Uint32(bytes[8:12])) |
||||
header.TargetNo = int32(binary.BigEndian.Uint32(bytes[12:16])) |
||||
cursor = 16 |
||||
default: |
||||
err = fmt.Errorf("unknown url type: %d", header.UrlType) |
||||
return |
||||
} |
||||
|
||||
payload = &Payload{Header: header} |
||||
payload.Body = bytes[cursor:] |
||||
return |
||||
} |
||||
|
||||
func Encode(payload *Payload) (bytes []byte, err error) { |
||||
header := payload.Header |
||||
headerLen := 16 |
||||
var svc, target []byte |
||||
var svcLen, targetLen, bodyLen int |
||||
if header.UrlType == 1 { |
||||
svc = []byte(header.Svc) |
||||
target = []byte(header.Target) |
||||
svcLen = len(svc) |
||||
targetLen = len(target) |
||||
headerLen += svcLen + targetLen |
||||
} |
||||
bodyLen = len(payload.Body) |
||||
bytes = make([]byte, headerLen+bodyLen) |
||||
bytes[0] = header.Magic |
||||
bytes[1] = header.Type |
||||
bytes[2] = header.Status |
||||
bytes[3] = (header.UrlType << 4) | header.SerializeType |
||||
binary.BigEndian.PutUint32(bytes[4:8], uint32(header.SeqId)) |
||||
|
||||
var cursor int |
||||
switch header.UrlType { |
||||
case 1: |
||||
binary.BigEndian.PutUint32(bytes[8:12], uint32(svcLen)) |
||||
copy(bytes[12:12+svcLen], svc) |
||||
cursor = 12 + svcLen |
||||
binary.BigEndian.PutUint32(bytes[cursor:cursor+4], uint32(targetLen)) |
||||
copy(bytes[cursor+4:cursor+4+targetLen], target) |
||||
cursor = cursor + 4 + targetLen |
||||
case 2: |
||||
binary.BigEndian.PutUint32(bytes[8:12], uint32(header.SvcNo)) |
||||
binary.BigEndian.PutUint32(bytes[12:16], uint32(header.TargetNo)) |
||||
cursor = 16 |
||||
default: |
||||
err = fmt.Errorf("unknown url type: %d", header.UrlType) |
||||
} |
||||
if bodyLen > 0 { |
||||
copy(bytes[cursor:], payload.Body) |
||||
} |
||||
return |
||||
} |
||||
|
||||
@ -0,0 +1,50 @@
|
||||
package protocol |
||||
|
||||
import ( |
||||
"errors" |
||||
"testing" |
||||
) |
||||
|
||||
func TestProtocolCodec(t *testing.T) { |
||||
header := Header{ |
||||
Magic: Magic, |
||||
Type: 1, |
||||
Status: 20, |
||||
UrlType: 1, |
||||
SerializeType: 2, |
||||
SeqId: 10086, |
||||
Svc: "Postal", |
||||
Target: "Deliver", |
||||
} |
||||
payload := &Payload{ |
||||
Header: header, |
||||
Body: []byte(`{"receiver":"10001啊"}`), |
||||
} |
||||
bytes, err := Encode(payload) |
||||
if err != nil { |
||||
t.Error(err) |
||||
return |
||||
} |
||||
payload2, err := Decode(bytes) |
||||
if err != nil { |
||||
t.Error(err) |
||||
return |
||||
} |
||||
header2 := payload2.Header |
||||
body2Len := len(payload2.Body) |
||||
ok := header2.Magic == header.Magic && |
||||
header2.Status == header.Status && |
||||
header2.UrlType == header.UrlType && |
||||
header2.SerializeType == header.SerializeType && |
||||
header2.SeqId == header.SeqId && |
||||
header2.SvcNo == header.SvcNo && |
||||
header2.TargetNo == header.TargetNo && |
||||
header2.Svc == header.Svc && |
||||
header2.Target == header.Target && |
||||
body2Len == len(payload.Body) && |
||||
payload2.Body[0] == payload.Body[0] && |
||||
payload2.Body[body2Len-1] == payload.Body[body2Len-1] |
||||
if !ok { |
||||
t.Error(errors.New("decode not equals")) |
||||
} |
||||
} |
||||
@ -0,0 +1,61 @@
|
||||
package resp |
||||
|
||||
import ( |
||||
"github.com/bytedance/sonic" |
||||
"github.com/cloudwego/kitex/pkg/klog" |
||||
) |
||||
|
||||
const ( |
||||
CodeOK = 200 |
||||
CodeFail = 400 |
||||
CodeError = 500 |
||||
) |
||||
|
||||
type H map[string]interface{} |
||||
|
||||
// Response 响应体包装
|
||||
type Response struct { |
||||
Seq int `json:"seq,omitempty"` |
||||
Code int `json:"code,omitempty"` |
||||
Msg string `json:"msg,omitempty"` |
||||
Data any `json:"data,omitempty"` |
||||
Extra map[string]any `json:"extra,omitempty"` |
||||
} |
||||
|
||||
func (resp *Response) Json() []byte { |
||||
json, err := sonic.Marshal(resp) |
||||
if err != nil { |
||||
klog.Error("unknown json error: ", err) |
||||
return nil |
||||
} |
||||
return json |
||||
} |
||||
|
||||
func (resp *Response) JsonString() string { |
||||
return string(resp.Json()) |
||||
} |
||||
|
||||
func SeqResp(seq int, code int, msg string, data any) *Response { |
||||
return &Response{ |
||||
Seq: seq, |
||||
Code: code, |
||||
Msg: msg, |
||||
Data: data, |
||||
} |
||||
} |
||||
|
||||
func Resp(code int, msg string, data any) *Response { |
||||
return SeqResp(0, code, msg, data) |
||||
} |
||||
|
||||
func Success(data any) *Response { |
||||
return Resp(CodeOK, "", data) |
||||
} |
||||
|
||||
func Fail(msg string) *Response { |
||||
return Resp(CodeFail, msg, nil) |
||||
} |
||||
|
||||
func Error(msg string) *Response { |
||||
return Resp(CodeError, msg, nil) |
||||
} |
||||
Loading…
Reference in new issue