Files
notify/msg.go
T

226 lines
5.8 KiB
Go
Raw Normal View History

2021-11-12 16:04:39 +08:00
package notify
import (
"context"
2021-11-12 16:04:39 +08:00
"net"
"time"
)
const defaultMessageReplyWriteTimeout = 30 * time.Second
2021-11-12 16:04:39 +08:00
const (
MSG_SYS MessageType = iota
MSG_SYS_WAIT
MSG_SYS_REPLY
// Deprecated: legacy RSA key-exchange control message.
2021-11-12 16:04:39 +08:00
MSG_KEY_CHANGE
MSG_ASYNC
MSG_SYNC_ASK
MSG_SYNC_REPLY
)
type MessageType uint8
type NetType uint8
const (
NET_SERVER NetType = iota
NET_CLIENT
)
type MsgVal []byte
type TransferMsg struct {
ID uint64
Key string
Value MsgVal
Type MessageType
}
type Message struct {
NetType
LogicalConn *LogicalConn
// Deprecated: ClientConn aliases LogicalConn for compatibility.
ClientConn *ClientConn
TransportConn *TransportConn
ServerConn Client
inboundTransportProfile *transportProtectionProfile
2021-11-12 16:04:39 +08:00
TransferMsg
Time time.Time
inboundConn net.Conn
2021-11-12 16:04:39 +08:00
}
type WaitMsg struct {
TransferMsg
Time time.Time
Reply chan Message
scope string
2021-11-12 16:04:39 +08:00
//Ctx context.Context
}
type messageLogicalTransferSender interface {
sendLogicalContext(context.Context, *LogicalConn, TransferMsg, time.Duration) (WaitMsg, error)
}
type messageTransportTransferSender interface {
sendTransportContextWithWriteTimeout(context.Context, *TransportConn, TransferMsg, time.Duration) (WaitMsg, error)
}
type messageInboundTransferSender interface {
sendTransferInboundContext(context.Context, *LogicalConn, *TransportConn, net.Conn, *transportProtectionProfile, TransferMsg, time.Duration) error
}
type messageClientTransferSender interface {
sendWithContextTimeout(context.Context, TransferMsg, time.Duration) (WaitMsg, error)
}
type messageReplyWriteTimeoutProvider interface {
ReplyWriteTimeout() time.Duration
}
2021-11-12 16:04:39 +08:00
func (m *Message) Reply(value MsgVal) (err error) {
return m.replyContext(context.Background(), value)
}
func (m *Message) ReplyCtx(ctx context.Context, value MsgVal) (err error) {
if ctx == nil {
ctx = context.Background()
}
return m.replyContext(ctx, value)
}
func (m *Message) replyContext(ctx context.Context, value MsgVal) (err error) {
logical := messageLogicalConnSnapshot(m)
transport := messageTransportConnSnapshot(m)
writeTimeout := defaultMessageReplyWriteTimeout
2021-11-12 16:04:39 +08:00
reply := TransferMsg{
ID: m.ID,
Key: m.Key,
Value: value,
Type: m.Type,
}
if reply.Type == MSG_SYNC_ASK {
reply.Type = MSG_SYNC_REPLY
}
if reply.Type == MSG_SYS_WAIT {
reply.Type = MSG_SYS_REPLY
}
if m.NetType == NET_SERVER {
if logical == nil {
return transportDetachedErrorForPeer(nil, transport)
}
server := logical.Server()
if server == nil {
return transportDetachedErrorForPeer(logical, transport)
}
if provider, ok := server.(messageReplyWriteTimeoutProvider); ok {
writeTimeout = provider.ReplyWriteTimeout()
}
if m.inboundConn != nil && logical != nil {
sender, _ := server.(messageInboundTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
return sender.sendTransferInboundContext(ctx, logical, transport, m.inboundConn, messageInboundTransportProtectionSnapshot(m), reply, writeTimeout)
}
if transport != nil {
sender, _ := server.(messageTransportTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
_, err = sender.sendTransportContextWithWriteTimeout(ctx, transport, reply, writeTimeout)
return err
}
sender, _ := server.(messageLogicalTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
_, err = sender.sendLogicalContext(ctx, logical, reply, writeTimeout)
2021-11-12 16:04:39 +08:00
}
if m.NetType == NET_CLIENT {
if m.ServerConn == nil {
return net.ErrClosed
}
if sender, ok := m.ServerConn.(messageClientTransferSender); ok {
_, err = sender.sendWithContextTimeout(ctx, reply, writeTimeout)
} else {
_, err = m.ServerConn.send(reply)
}
2021-11-12 16:04:39 +08:00
}
return
}
func (m *Message) ReplyObj(value interface{}) (err error) {
return m.ReplyObjCtx(context.Background(), value)
}
func (m *Message) ReplyObjCtx(ctx context.Context, value interface{}) (err error) {
2021-11-12 16:04:39 +08:00
data, err := encode(value)
if err != nil {
return err
}
return m.ReplyCtx(ctx, data)
2021-11-12 16:04:39 +08:00
}
func hydrateServerMessagePeerFields(message Message) Message {
if message.LogicalConn == nil {
message.LogicalConn = logicalConnFromClient(message.ClientConn)
2021-11-12 16:04:39 +08:00
}
if message.LogicalConn == nil && message.TransportConn != nil {
message.LogicalConn = message.TransportConn.logicalConnSnapshot()
}
if message.ClientConn == nil && message.LogicalConn != nil {
message.ClientConn = message.LogicalConn.compatClientConn()
2021-11-12 16:04:39 +08:00
}
if message.TransportConn == nil && message.LogicalConn != nil {
message.TransportConn = message.LogicalConn.CurrentTransportConn()
2021-11-12 16:04:39 +08:00
}
if message.inboundConn != nil && message.inboundTransportProfile == nil && message.LogicalConn != nil {
profile := message.LogicalConn.transportProtectionProfileSnapshot()
message.inboundTransportProfile = &profile
}
return message
2021-11-12 16:04:39 +08:00
}
func messageLogicalConnSnapshot(message *Message) *LogicalConn {
if message == nil {
return nil
2021-11-12 16:04:39 +08:00
}
if message.LogicalConn != nil {
return message.LogicalConn
2021-11-12 16:04:39 +08:00
}
return logicalConnFromClient(message.ClientConn)
2021-11-12 16:04:39 +08:00
}
func messageTransportConnSnapshot(message *Message) *TransportConn {
if message == nil {
return nil
2022-05-19 11:04:52 +08:00
}
if message.TransportConn != nil {
return message.TransportConn
2022-05-19 11:04:52 +08:00
}
logical := messageLogicalConnSnapshot(message)
if logical == nil {
return nil
2022-05-19 11:04:52 +08:00
}
return logical.CurrentTransportConn()
2022-05-19 11:04:52 +08:00
}
func messageInboundTransportProtectionSnapshot(message *Message) *transportProtectionProfile {
if message == nil {
return nil
}
if message.inboundTransportProfile != nil {
return message.inboundTransportProfile
}
if message.inboundConn == nil {
return nil
}
logical := messageLogicalConnSnapshot(message)
if logical == nil {
return nil
}
profile := logical.transportProtectionProfileSnapshot()
message.inboundTransportProfile = &profile
return message.inboundTransportProfile
}