Files
anytunnel/at-common/message-channel.go
2019-08-08 17:13:34 +08:00

316 lines
7.6 KiB
Go

package at_common
import (
"bufio"
"bytes"
"crypto/tls"
"encoding/binary"
"encoding/json"
"fmt"
"math/rand"
"net"
"reflect"
"sync"
"time"
)
type Msg struct {
MsgType int
}
type message struct {
Type int
Data interface{}
}
type msgErrorHandler func(channel *MessageChannel, rawMsg interface{}, err error)
type msgCallback func(channel *MessageChannel, msg interface{})
type closeCallback func(channel *MessageChannel, isPeerClose bool)
type MessageChannel struct {
reader *bufio.Reader
writer *bufio.Writer
msgHandler map[int][]msgCallback
msgErrorHandler msgErrorHandler
msgTypeMap map[int]interface{}
readlock *sync.Mutex
writelock *sync.Mutex
readHandler func(msg message)
serveStared bool
ConnectionID uint64
Conn *net.Conn
remoteAddr net.Addr
localAddr net.Addr
}
func NewMessageChannel(conn *net.Conn) MessageChannel {
ch := MessageChannel{
Conn: conn,
remoteAddr: (*conn).RemoteAddr(),
localAddr: (*conn).LocalAddr(),
reader: bufio.NewReader(*conn),
writer: bufio.NewWriter(*conn),
msgTypeMap: map[int]interface{}{},
readlock: &sync.Mutex{},
writelock: &sync.Mutex{},
ConnectionID: rand.Uint64(),
msgErrorHandler: func(channel *MessageChannel, rawMsg interface{}, err error) {},
msgHandler: map[int][]msgCallback{},
}
ch.RegMsg(MSG_TYPE_PING, new(MsgPing), func(channel *MessageChannel, msg interface{}) {
msgPing := msg.(*MsgPing)
channel.Pong(msgPing.ID)
return
})
return ch
}
func NewMessageChannelTls(conn *tls.Conn) MessageChannel {
con := net.Conn(conn)
return NewMessageChannel(&con)
}
func (mc *MessageChannel) CloseConn() (err error) {
(*mc.Conn).SetDeadline(time.Now().Add(time.Millisecond))
return (*mc.Conn).Close()
}
func (mc *MessageChannel) RegMsg(msgType int, msg interface{}, fn msgCallback, fns ...msgCallback) {
mc.msgTypeMap[msgType] = msg
mc.msgHandler[msgType] = append(fns, fn)
}
func (mc *MessageChannel) SetMsgErrorHandler(fn func(channel *MessageChannel, rawMsg interface{}, err error)) {
mc.msgErrorHandler = fn
}
func (mc *MessageChannel) encode(data interface{}) (msg []byte, err error) {
var message []byte
message, err = json.Marshal(data)
if err != nil {
return
}
// 读取消息的长度
var length int32 = int32(len(message))
var pkg *bytes.Buffer = new(bytes.Buffer)
// 写入消息头
err = binary.Write(pkg, binary.LittleEndian, length)
if err != nil {
return
}
// 写入消息实体
err = binary.Write(pkg, binary.LittleEndian, message)
if err != nil {
return
}
return pkg.Bytes(), nil
}
func (mc *MessageChannel) Write(msg interface{}) (err error) {
defer mc.writelock.Unlock()
mc.writelock.Lock()
if reflect.TypeOf(msg).Kind().String() != "struct" {
err = fmt.Errorf("error : message must be a struct , send to %s", mc.remoteAddr)
return
}
if _, ok := reflect.TypeOf(msg).FieldByName("MsgType"); !ok {
err = fmt.Errorf("error : message must be has MsgType field , send to %s", mc.remoteAddr)
return
}
if reflect.ValueOf(msg).FieldByName("MsgType").Int() <= 0 {
err = fmt.Errorf("error : message's MsgType field must be great than 0, send to %s", mc.remoteAddr)
return
}
v := reflect.ValueOf(msg).FieldByName("MsgType").Int()
pack := message{
Type: int(v),
Data: msg,
}
var msgData []byte
msgData, err = mc.encode(pack)
if err != nil {
err = fmt.Errorf("encode message error : %s , to : %s", err, mc.remoteAddr)
return
}
_, err = mc.writer.Write(msgData)
if err != nil {
err = fmt.Errorf("write messasge fail to %s , ERR : %s", mc.remoteAddr, err)
return
}
err = mc.writer.Flush()
if err != nil {
err = fmt.Errorf("flush messasge fail %s , ERR : %s", mc.remoteAddr, err)
return
}
return
}
func (mc *MessageChannel) ReadTimeout(msg interface{}, timeout int) (err error) {
if !mc.serveStared {
err = fmt.Errorf("DoServe() must be called before call Read(),From : %s", mc.remoteAddr)
return
}
type M struct {
Err error
Msg interface{}
}
msgChn := make(chan M, 1)
mc.readHandler = func(rawMsg message) {
_, err = mc.toStruct(rawMsg.Data, &msg)
msgChn <- M{
Msg: msg,
Err: err,
}
}
m := M{}
if timeout > 0 {
select {
case m = <-msgChn:
msg = m.Msg
err = m.Err
case <-time.After(time.Duration(timeout) * time.Millisecond):
err = fmt.Errorf("read channel message timeout from %s", mc.remoteAddr)
}
} else {
m = <-msgChn
msg = m.Msg
err = m.Err
}
return
}
func (mc *MessageChannel) Read(msg interface{}) (err error) {
return mc.ReadTimeout(msg, 0)
}
func (mc *MessageChannel) read() (msg message, err error) {
defer func() {
if err != nil {
(*mc.Conn).Close()
}
mc.readlock.Unlock()
}()
mc.readlock.Lock()
// 读取消息的长度
lengthByte, err := mc.reader.Peek(4)
if err != nil {
err = fmt.Errorf("read message length error : %s , from %s", err, mc.remoteAddr)
return
}
lengthBuff := bytes.NewBuffer(lengthByte)
var length int32
err = binary.Read(lengthBuff, binary.LittleEndian, &length)
if err != nil {
err = fmt.Errorf("read message error : %s , from %s", err, mc.remoteAddr)
return
}
if int32(mc.reader.Buffered()) < length+4 {
err = fmt.Errorf("message data length error from %s", mc.remoteAddr)
return
}
// 读取消息真正的内容
pack := make([]byte, int(4+length))
_, err = mc.reader.Read(pack)
if err != nil {
err = fmt.Errorf("read message error : %s , from %s", err, mc.remoteAddr)
return
}
err = json.Unmarshal(pack[4:], &msg)
if err != nil {
err = fmt.Errorf("unmarshal message error : %s , from %s", err, mc.remoteAddr)
return
}
return
}
func (mc *MessageChannel) DoServe(errfn func(err error)) {
mc.serveStared = true
go func() {
var err error
var msg message
for {
msg, err = mc.read()
if err != nil {
go errfn(err)
mc.serveStared = false
return
}
if mc.readHandler != nil {
go mc.readHandler(msg)
mc.readHandler = nil
continue
}
h, ok := mc.msgHandler[msg.Type]
var data interface{}
if ok {
data, err = mc.parse(msg)
} else {
err = fmt.Errorf("msg handler not found , msgType:%d", msg.Type)
}
if err != nil {
go mc.msgErrorHandler(mc, msg.Data, err)
} else {
for i := len(h) - 1; i >= 0; i-- {
go h[i](mc, data)
}
}
}
}()
return
}
func (mc *MessageChannel) Ping() (id string, err error) {
id = randStr(32)
ping := MsgPing{
Msg: Msg{MsgType: MSG_TYPE_PING},
ID: id,
}
err = mc.Write(ping)
return
}
func (mc *MessageChannel) Pong(id string) (err error) {
pong := MsgPong{
Msg: Msg{MsgType: MSG_TYPE_PONG},
ID: id,
}
err = mc.Write(pong)
return
}
func (mc *MessageChannel) RemoteAddr() net.Addr {
return mc.remoteAddr
}
func (mc *MessageChannel) LocalAddr() net.Addr {
return mc.localAddr
}
func (mc *MessageChannel) parse(msg message) (data interface{}, err error) {
data, ok := mc.msgTypeMap[msg.Type]
if !ok {
err = fmt.Errorf("message type not registed")
return
}
mbytes, err := json.Marshal(msg.Data)
if err != nil {
return
}
err = json.Unmarshal(mbytes, data)
return
}
func (mc *MessageChannel) toStruct(msg, struc interface{}) (data interface{}, err error) {
mbytes, err := json.Marshal(msg)
if err != nil {
return
}
err = json.Unmarshal(mbytes, &struc)
data = struc
return
}
func randStr(strlen int) string {
codes := "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
codeLen := len(codes)
data := make([]byte, strlen)
rand.Seed(time.Now().UnixNano() + rand.Int63() + rand.Int63() + rand.Int63() + rand.Int63())
for i := 0; i < strlen; i++ {
idx := rand.Intn(codeLen)
data[i] = byte(codes[idx])
}
return string(data)
}