proto-socket/go/tcp_client.go
toki cab031ea28 Refactor client implementations with baseClient pattern
- Dart: Add explicit Future<void> return type to send() method
- Dart: Improve onDisconnected() logic with proper isAlive handling and heartbeat response
- Go: Introduce baseClient generic base class for shared client logic
- Go: Refactor TcpClient and WsClient to embed baseClient instead of embedding Communicator
- Remove duplicate heartbeat and disconnect handling code across client implementations
- Clean up unused imports (time package)
2026-04-11 13:52:50 +09:00

120 lines
2.7 KiB
Go

package toki_socket
import (
"context"
"crypto/tls"
"encoding/binary"
"fmt"
"io"
"net"
"sync"
"google.golang.org/protobuf/proto"
"toki-labs.com/toki_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
}