package proto_socket import ( "context" "crypto/tls" "errors" "fmt" "net" "sync" "google.golang.org/protobuf/proto" ) type TcpServer struct { host string port int tlsConfig *tls.Config newClient func(net.Conn) *TcpClient clients []*TcpClient mu sync.Mutex listener net.Listener started bool OnClientConnected func(*TcpClient) } func NewTcpServer(host string, port int, newClient func(net.Conn) *TcpClient) *TcpServer { return &TcpServer{ host: host, port: port, newClient: newClient, OnClientConnected: func(*TcpClient) {}, } } func NewTcpServerTLS(host string, port int, tlsCfg *tls.Config, newClient func(net.Conn) *TcpClient) *TcpServer { s := NewTcpServer(host, port, newClient) s.tlsConfig = tlsCfg return s } func (s *TcpServer) Started() bool { s.mu.Lock() defer s.mu.Unlock() return s.started } func (s *TcpServer) IsSecure() bool { return s.tlsConfig != nil } func (s *TcpServer) Clients() []*TcpClient { s.mu.Lock() defer s.mu.Unlock() return append([]*TcpClient{}, s.clients...) } func (s *TcpServer) Start(ctx context.Context) error { addr := fmt.Sprintf("%s:%d", s.host, s.port) ln, err := net.Listen("tcp", addr) if err != nil { return err } if s.tlsConfig != nil { ln = tls.NewListener(ln, s.tlsConfig) } s.mu.Lock() s.listener = ln s.started = true s.mu.Unlock() go func() { <-ctx.Done() _ = s.Stop() }() go s.acceptLoop() return nil } func (s *TcpServer) Stop() error { s.mu.Lock() if !s.started { s.mu.Unlock() return nil } s.started = false ln := s.listener clients := append([]*TcpClient{}, s.clients...) s.clients = nil s.mu.Unlock() var err error if ln != nil { err = ln.Close() } for _, client := range clients { _ = client.Close() } return err } func (s *TcpServer) Broadcast(m proto.Message) error { clients := s.Clients() var errs []error for _, client := range clients { if err := client.Send(m); err != nil { errs = append(errs, err) } } return errors.Join(errs...) } func (s *TcpServer) acceptLoop() { for { conn, err := s.listener.Accept() if err != nil { s.mu.Lock() started := s.started s.mu.Unlock() if !started || errors.Is(err, net.ErrClosed) { return } continue } client := s.newClient(conn) client.AddDisconnectListener(func(c *TcpClient) { s.removeClient(c) }) s.mu.Lock() s.clients = append(s.clients, client) onConnected := s.OnClientConnected s.mu.Unlock() onConnected(client) } } func (s *TcpServer) removeClient(client *TcpClient) { s.mu.Lock() defer s.mu.Unlock() for i, item := range s.clients { if item == client { s.clients = append(s.clients[:i], s.clients[i+1:]...) return } } }