120 lines
2.7 KiB
Go
120 lines
2.7 KiB
Go
package toki_socket
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"encoding/binary"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"sync"
|
|
|
|
"google.golang.org/protobuf/proto"
|
|
|
|
"git.toki-labs.com/toki/common-proto-socket/go/packets"
|
|
)
|
|
|
|
const MaxPacketSize = 64 << 20
|
|
|
|
type TcpClient struct {
|
|
baseClient[*TcpClient]
|
|
conn net.Conn
|
|
writeMu sync.Mutex // Defensive if WritePacket is called outside the queued writer.
|
|
}
|
|
|
|
func NewTcpClient(conn net.Conn, intervalSec, waitSec int, parserMap ParserMap) *TcpClient {
|
|
c := &TcpClient{
|
|
conn: conn,
|
|
}
|
|
c.baseClient = newBaseClient(c, intervalSec, waitSec, func() error {
|
|
return c.conn.Close()
|
|
})
|
|
c.Communicator.Initialize(c, parserMap)
|
|
c.SetWriteErrorHandler(func(error) {
|
|
c.onDisconnected()
|
|
})
|
|
c.AddListener(TypeNameOf(&packets.HeartBeat{}), func(m proto.Message) {
|
|
c.onHeartBeat()
|
|
})
|
|
go c.readLoop()
|
|
c.sendHeartBeat()
|
|
return c
|
|
}
|
|
|
|
func DialTcp(ctx context.Context, host string, port int, intervalSec, waitSec int, parserMap ParserMap) (*TcpClient, error) {
|
|
var d net.Dialer
|
|
conn, err := d.DialContext(ctx, "tcp", fmt.Sprintf("%s:%d", host, port))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return NewTcpClient(conn, intervalSec, waitSec, parserMap), nil
|
|
}
|
|
|
|
func DialTcpTLS(ctx context.Context, host string, port int, tlsCfg *tls.Config, intervalSec, waitSec int, parserMap ParserMap) (*TcpClient, error) {
|
|
d := tls.Dialer{Config: tlsCfg}
|
|
conn, err := d.DialContext(ctx, "tcp", fmt.Sprintf("%s:%d", host, port))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return NewTcpClient(conn, intervalSec, waitSec, parserMap), nil
|
|
}
|
|
|
|
func (c *TcpClient) readLoop() {
|
|
header := make([]byte, 4)
|
|
for c.IsAlive() {
|
|
if _, err := io.ReadFull(c.conn, header); err != nil {
|
|
c.onDisconnected()
|
|
return
|
|
}
|
|
length := binary.BigEndian.Uint32(header)
|
|
if length == 0 {
|
|
continue
|
|
}
|
|
if length > MaxPacketSize {
|
|
c.onDisconnected()
|
|
return
|
|
}
|
|
packetBytes := make([]byte, int(length))
|
|
if _, err := io.ReadFull(c.conn, packetBytes); err != nil {
|
|
c.onDisconnected()
|
|
return
|
|
}
|
|
base := &packets.PacketBase{}
|
|
if err := proto.Unmarshal(packetBytes, base); err != nil {
|
|
c.onDisconnected()
|
|
return
|
|
}
|
|
c.OnReceivedData(base.GetTypeName(), base.GetData(), base.GetNonce(), base.GetResponseNonce())
|
|
c.sendHeartBeat()
|
|
}
|
|
}
|
|
|
|
func (c *TcpClient) WritePacket(base *packets.PacketBase) error {
|
|
b, err := proto.Marshal(base)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
header := make([]byte, 4)
|
|
binary.BigEndian.PutUint32(header, uint32(len(b)))
|
|
|
|
c.writeMu.Lock()
|
|
defer c.writeMu.Unlock()
|
|
if err := writeFull(c.conn, header); err != nil {
|
|
return err
|
|
}
|
|
return writeFull(c.conn, b)
|
|
}
|
|
|
|
func writeFull(w io.Writer, b []byte) error {
|
|
for len(b) > 0 {
|
|
n, err := w.Write(b)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if n == 0 {
|
|
return io.ErrShortWrite
|
|
}
|
|
b = b[n:]
|
|
}
|
|
return nil
|
|
}
|