package transport import ( "context" "crypto/tls" "errors" "fmt" "net" "sync" "sync/atomic" "go.uber.org/zap" ) // Server is the IOP TCP transport server. type Server struct { handler Handler logger *zap.Logger tlsConfig *tls.Config listener net.Listener started atomic.Bool mu sync.Mutex sessions map[string]*Session nextID atomic.Uint64 } func NewServer(handler Handler, logger *zap.Logger) *Server { return &Server{ handler: handler, logger: logger, sessions: make(map[string]*Session), } } // SetTLS configures mTLS for the server. Must be called before Start. func (s *Server) SetTLS(cfg *tls.Config) { s.tlsConfig = cfg } // Start begins accepting connections on addr. func (s *Server) Start(ctx context.Context, addr string) error { ln, err := net.Listen("tcp", addr) if err != nil { return fmt.Errorf("transport: listen %s: %w", addr, err) } if s.tlsConfig != nil { ln = tls.NewListener(ln, s.tlsConfig) } s.mu.Lock() s.listener = ln s.started.Store(true) s.mu.Unlock() s.logger.Info("transport listening", zap.String("addr", addr), zap.Bool("tls", s.tlsConfig != nil), ) go func() { <-ctx.Done() _ = s.Stop() }() go s.acceptLoop(ctx) return nil } // Stop closes the listener and all active sessions. func (s *Server) Stop() error { if !s.started.CompareAndSwap(true, false) { return nil } s.mu.Lock() ln := s.listener sessions := make([]*Session, 0, len(s.sessions)) for _, sess := range s.sessions { sessions = append(sessions, sess) } s.mu.Unlock() var err error if ln != nil { err = ln.Close() } for _, sess := range sessions { sess.Close() } s.logger.Info("transport stopped") return err } // SessionCount returns the number of active sessions. func (s *Server) SessionCount() int { s.mu.Lock() defer s.mu.Unlock() return len(s.sessions) } func (s *Server) acceptLoop(ctx context.Context) { for { conn, err := s.listener.Accept() if err != nil { if !s.started.Load() || errors.Is(err, net.ErrClosed) { return } s.logger.Warn("accept error", zap.Error(err)) continue } go s.handleConn(ctx, conn) } } func (s *Server) handleConn(ctx context.Context, conn net.Conn) { id := fmt.Sprintf("sess-%d", s.nextID.Add(1)) sess := newSession(id, conn, s.handler, s.logger) s.mu.Lock() s.sessions[id] = sess s.mu.Unlock() s.logger.Info("client connected", zap.String("session_id", id), zap.String("remote", conn.RemoteAddr().String()), ) sess.readLoop(ctx) s.mu.Lock() delete(s.sessions, id) s.mu.Unlock() s.logger.Info("client disconnected", zap.String("session_id", id)) }