package proto_socket import ( "context" "crypto/tls" "encoding/binary" "errors" "fmt" "io" "net" "sync" "google.golang.org/protobuf/proto" "git.toki-labs.com/toki/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(err error) { c.onDisconnected(DisconnectReasonWriteError, err) }) 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(tcpReadDisconnectReason(DisconnectReasonTCPReadHeader, err), err) return } length := binary.BigEndian.Uint32(header) if length == 0 { continue } if length > MaxPacketSize { c.onDisconnected(DisconnectReasonTCPPacketTooLarge, fmt.Errorf("packet size %d exceeds max %d", length, MaxPacketSize)) return } packetBytes := make([]byte, int(length)) if _, err := io.ReadFull(c.conn, packetBytes); err != nil { c.onDisconnected(tcpReadDisconnectReason(DisconnectReasonTCPReadPayload, err), err) return } base := &packets.PacketBase{} if err := proto.Unmarshal(packetBytes, base); err != nil { c.onDisconnected(DisconnectReasonTCPPacketParse, err) 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 } func tcpReadDisconnectReason(fallback string, err error) string { if errors.Is(err, io.EOF) { return DisconnectReasonRemoteClosed } return fallback }