Files
2023-02-20 15:04:50 +08:00

333 lines
7.5 KiB
Go

package conn
import (
"context"
"crypto/tls"
"net"
"strings"
"sync"
"time"
"github.com/lwch/logging"
"github.com/lwch/natpass/code/client/global"
"github.com/lwch/natpass/code/network"
"github.com/lwch/natpass/code/utils"
"github.com/lwch/runtime"
)
const dropBlockTimeout = 10 * time.Minute
// Conn connection
type Conn struct {
sync.RWMutex
cfg *global.Configure // configure
conn *network.Conn // connection wrap, read write with timeout
read map[string]chan *network.Msg // link id => channel
unknownRead chan *network.Msg // read message without link
onDisconnect chan string // on disconnect channel, the value is clientid
write chan *network.Msg // write queue
lockDrop sync.RWMutex // drop mutex
drop map[string]time.Time // drop cache, drop message when this link is closed
// runtime
ctx context.Context
cancel context.CancelFunc
}
// New new connection
func New(cfg *global.Configure) *Conn {
conn := &Conn{
cfg: cfg,
read: make(map[string]chan *network.Msg),
unknownRead: make(chan *network.Msg, 1024),
onDisconnect: make(chan string, 1024),
write: make(chan *network.Msg, 10*1024*1024),
drop: make(map[string]time.Time),
}
runtime.Assert(conn.connect())
conn.ctx, conn.cancel = context.WithCancel(context.Background())
go conn.loopRead()
go conn.loopWrite()
go conn.keepalive()
go conn.checkDrop()
return conn
}
// connect connect server and write handshake packet
func (conn *Conn) connect() error {
var dial net.Conn
var err error
if conn.cfg.UseSSL {
if conn.cfg.SSLInsecure { // disable sni
rawConn, err := net.Dial("tcp", conn.cfg.Server)
if err != nil {
logging.Error("raw dial: %v", err)
return err
}
cfg := new(tls.Config)
cfg.InsecureSkipVerify = true
dial = tls.Client(rawConn, cfg)
err = dial.(*tls.Conn).Handshake()
if err != nil {
rawConn.Close()
}
} else {
dial, err = tls.Dial("tcp", conn.cfg.Server, nil)
}
} else {
dial, err = net.Dial("tcp", conn.cfg.Server)
}
if err != nil {
logging.Error("dial: %v", err)
return err
}
cn := network.NewConn(dial)
err = writeHandshake(cn, conn.cfg)
if err != nil {
logging.Error("write handshake: %v", err)
return err
}
logging.Info("%s connected", conn.cfg.Server)
conn.conn = cn
return nil
}
func (conn *Conn) close() {
if conn.conn != nil {
conn.conn.Close()
}
}
// writeHandshake send handshake message, default timeout is 5 seconds
func writeHandshake(conn *network.Conn, cfg *global.Configure) error {
var msg network.Msg
msg.XType = network.Msg_handshake
msg.From = cfg.ID
msg.To = "server"
msg.Payload = &network.Msg_Hsp{
Hsp: &network.HandshakePayload{
Enc: cfg.Hasher.Hash(),
},
}
return conn.WriteMessage(&msg, 5*time.Second)
}
// isDrop check the message is dropped by linkid
func (conn *Conn) isDrop(linkID string) bool {
conn.lockDrop.RLock()
defer conn.lockDrop.RUnlock()
_, ok := conn.drop[linkID]
return ok
}
// addDrop add to drop queue
func (conn *Conn) addDrop(linkID string) {
conn.lockDrop.Lock()
defer conn.lockDrop.Unlock()
conn.drop[linkID] = time.Now().Add(dropBlockTimeout)
}
// getChan get read channel by linkid
func (conn *Conn) getChan(linkID string) chan *network.Msg {
conn.RLock()
ch := conn.read[linkID]
conn.RUnlock()
if ch == nil {
ch = conn.unknownRead
}
return ch
}
// hookDispatch hook message before dispatcher
func (conn *Conn) hookDispatch(msg *network.Msg) bool {
switch msg.GetXType() {
// if disconnected add linkid to drop list, and break the handle chain
case network.Msg_disconnect:
conn.addDrop(msg.GetLinkId())
// TODO: no need will block
// conn.onDisconnect <- msg.GetLinkId()
logging.Info("connection %s disconnected", msg.GetLinkId())
return false
}
return true
}
// handleLinkedMessage linked message handler, return false means break read loop
func (conn *Conn) handleLinkedMessage(msg *network.Msg) bool {
linkID := msg.GetLinkId()
if conn.isDrop(linkID) {
return true
}
if !conn.hookDispatch(msg) {
return true
}
ch := conn.getChan(linkID)
select {
case ch <- msg:
case <-time.After(conn.cfg.WriteTimeout):
logging.Error("drop message: %s", msg.GetXType().String())
conn.addDrop(linkID)
case <-conn.ctx.Done():
return false
}
return true
}
// handleUnlinkedMessage unlinked message handler, return false means break read loop
func (conn *Conn) handleUnlinkedMessage(msg *network.Msg) bool {
// TODO
return true
}
// loopRead loop read message
func (conn *Conn) loopRead() {
defer utils.Recover("loopRead")
defer conn.close()
defer conn.cancel()
var timeout int
run := func(msg *network.Msg) bool {
timeout = 0
// skip keepalive message
if msg.GetXType() == network.Msg_keepalive {
return true
}
logging.Debug("read message %s(%s) from %s",
msg.GetXType().String(), msg.GetLinkId(), msg.GetFrom())
linkID := msg.GetLinkId()
if len(linkID) > 0 {
return conn.handleLinkedMessage(msg)
}
return conn.handleUnlinkedMessage(msg)
}
for {
msg, _, err := conn.conn.ReadMessage(conn.cfg.ReadTimeout)
if err != nil {
if strings.Contains(err.Error(), "i/o timeout") {
timeout++
if timeout >= 60 {
logging.Error("too many timeout times")
return
}
continue
}
logging.Error("read message: %v", err)
return
}
if !run(msg) {
return
}
}
}
// loopWrite loop write message
func (conn *Conn) loopWrite() {
defer utils.Recover("loopWrite")
defer conn.close()
defer conn.cancel()
for {
var msg *network.Msg
select {
case msg = <-conn.write:
case <-conn.ctx.Done():
return
}
msg.From = conn.cfg.ID
err := conn.conn.WriteMessage(msg, conn.cfg.WriteTimeout)
if err != nil {
logging.Error("write message error on %s: %v",
conn.cfg.ID, err)
continue
}
}
}
// keepalive loop send keepalive message
func (conn *Conn) keepalive() {
defer utils.Recover("keepalive")
defer conn.close()
defer conn.cancel()
tk := time.NewTicker(10 * time.Second)
for {
select {
case <-tk.C:
conn.SendKeepalive()
case <-conn.ctx.Done():
return
}
}
}
// AddLink attach read message
func (conn *Conn) AddLink(id string) {
logging.Info("add link %s", id)
conn.Lock()
if _, ok := conn.read[id]; !ok {
conn.read[id] = make(chan *network.Msg, 1024)
}
conn.Unlock()
}
// Requeue requeue for next read
func (conn *Conn) Requeue(id string, msg *network.Msg) {
conn.RLock()
ch := conn.read[id]
conn.RUnlock()
ch <- msg
}
// ChanRead get read channel from link id
func (conn *Conn) ChanRead(id string) <-chan *network.Msg {
conn.RLock()
defer conn.RUnlock()
return conn.read[id]
}
// ChanUnknown get channel of unknown link id
func (conn *Conn) ChanUnknown() <-chan *network.Msg {
return conn.unknownRead
}
// ChanDisconnect get channel of disconnect
func (conn *Conn) ChanDisconnect() <-chan string {
return conn.onDisconnect
}
// checkDrop clear timeouted drop queue
func (conn *Conn) checkDrop() {
for {
time.Sleep(time.Second)
drops := make([]string, 0, len(conn.drop))
conn.lockDrop.RLock()
for k, t := range conn.drop {
if time.Now().After(t) {
drops = append(drops, k)
}
}
conn.lockDrop.RUnlock()
conn.lockDrop.Lock()
for _, id := range drops {
delete(conn.drop, id)
}
conn.lockDrop.Unlock()
}
}
// Wait wait for connection closed
func (conn *Conn) Wait() {
<-conn.ctx.Done()
}
// ChanClose close read chan
func (conn *Conn) ChanClose(id string) {
conn.Lock()
ch := conn.read[id]
if ch != nil {
close(ch)
}
delete(conn.read, id)
conn.Unlock()
conn.addDrop(id)
}