Add close message type

This commit is contained in:
Zoltan Papp
2024-06-05 19:49:30 +02:00
parent a40d4d2f32
commit fed9e587af
14 changed files with 371 additions and 170 deletions
+1
View File
@@ -5,4 +5,5 @@ import "net"
type Listener interface {
Listen(func(conn net.Conn)) error
Close() error
WaitForExitAcceptedConns()
}
+6 -1
View File
@@ -21,6 +21,11 @@ type Listener struct {
lock sync.Mutex
}
func (l *Listener) WaitForExitAcceptedConns() {
l.wg.Wait()
return
}
func NewListener(address string) listener.Listener {
return &Listener{
address: address,
@@ -61,11 +66,11 @@ func (l *Listener) Close() error {
l.lock.Lock()
defer l.lock.Unlock()
log.Infof("closing UDP server")
if l.listener == nil {
return nil
}
log.Infof("closing UDP listener")
close(l.quit)
err := l.listener.Close()
l.wg.Wait()
+8 -1
View File
@@ -33,8 +33,11 @@ func NewListener(address string) listener.Listener {
}
}
// Listen todo: prevent multiple call
func (l *Listener) Listen(acceptFn func(conn net.Conn)) error {
if l.server != nil {
return errors.New("server is already running")
}
l.acceptFn = acceptFn
mux := http.NewServeMux()
mux.HandleFunc("/", l.onAccept)
@@ -69,6 +72,10 @@ func (l *Listener) Close() error {
return nil
}
func (l *Listener) WaitForExitAcceptedConns() {
l.wg.Wait()
}
func (l *Listener) onAccept(writer http.ResponseWriter, request *http.Request) {
l.wg.Add(1)
defer l.wg.Done()
+21 -6
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"io"
"net"
"sync"
"time"
log "github.com/sirupsen/logrus"
@@ -17,7 +18,9 @@ type Conn struct {
lAddr *net.TCPAddr
rAddr *net.TCPAddr
ctx context.Context
closed bool
closedMu sync.Mutex
ctx context.Context
}
func NewConn(wsConn *websocket.Conn, lAddr, rAddr *net.TCPAddr) *Conn {
@@ -32,7 +35,7 @@ func NewConn(wsConn *websocket.Conn, lAddr, rAddr *net.TCPAddr) *Conn {
func (c *Conn) Read(b []byte) (n int, err error) {
t, r, err := c.Reader(c.ctx)
if err != nil {
return 0, ioErrHandling(err)
return 0, c.ioErrHandling(err)
}
if t != websocket.MessageBinary {
@@ -42,7 +45,7 @@ func (c *Conn) Read(b []byte) (n int, err error) {
n, err = r.Read(b)
if err != nil {
return 0, ioErrHandling(err)
return 0, c.ioErrHandling(err)
}
return n, err
}
@@ -76,11 +79,23 @@ func (c *Conn) SetDeadline(t time.Time) error {
}
func (c *Conn) Close() error {
return c.Conn.Close(websocket.StatusNormalClosure, "")
c.closedMu.Lock()
c.closed = true
c.closedMu.Unlock()
return c.Conn.CloseNow()
}
// todo: fix io.EOF handling
func ioErrHandling(err error) error {
func (c *Conn) isClosed() bool {
c.closedMu.Lock()
defer c.closedMu.Unlock()
return c.closed
}
func (c *Conn) ioErrHandling(err error) error {
if c.isClosed() {
return io.EOF
}
var wErr *websocket.CloseError
if !errors.As(err, &wErr) {
return err
+6 -3
View File
@@ -56,15 +56,18 @@ func (l *Listener) Close() error {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
log.Debugf("closing WS server")
log.Infof("stop WS listener")
if err := l.server.Shutdown(ctx); err != nil {
return fmt.Errorf("server shutdown failed: %v", err)
}
l.wg.Wait()
log.Infof("WS listener stopped")
return nil
}
func (l *Listener) WaitForExitAcceptedConns() {
l.wg.Wait()
}
func (l *Listener) onAccept(w http.ResponseWriter, r *http.Request) {
l.wg.Add(1)
defer l.wg.Done()
+33 -6
View File
@@ -15,11 +15,9 @@ import (
ws "github.com/netbirdio/netbird/relay/server/listener/wsnhooyr"
)
// Server
// todo:
// authentication: provide JWT token via RPC call. The MGM server can forward the token to the agents.
type Server struct {
store *Store
store *Store
storeMu sync.RWMutex
UDPListener listener.Listener
WSListener listener.Listener
@@ -27,7 +25,8 @@ type Server struct {
func NewServer() *Server {
return &Server{
store: NewStore(),
store: NewStore(),
storeMu: sync.RWMutex{},
}
}
@@ -69,6 +68,11 @@ func (r *Server) Close() error {
if r.UDPListener != nil {
uErr = r.UDPListener.Close()
}
r.sendCloseMsgs()
r.WSListener.WaitForExitAcceptedConns()
err := errors.Join(wErr, uErr)
return err
}
@@ -88,7 +92,7 @@ func (r *Server) accept(conn net.Conn) {
r.store.AddPeer(peer)
defer func() {
r.store.DeletePeer(peer)
peer.Log.Infof("peer left")
peer.Log.Infof("relay connection closed")
}()
for {
@@ -132,10 +136,33 @@ func (r *Server) accept(conn net.Conn) {
}
return
}()
case messages.MsgClose:
peer.Log.Infof("peer disconnected gracefully")
_ = conn.Close()
return
}
}
}
func (r *Server) sendCloseMsgs() {
msg := messages.MarshalCloseMsg()
r.storeMu.Lock()
log.Debugf("sending close messages to %d peers", len(r.store.peers))
for _, p := range r.store.peers {
_, err := p.conn.Write(msg)
if err != nil {
log.Errorf("failed to send close message to peer: %s", p.String())
}
err = p.conn.Close()
if err != nil {
log.Errorf("failed to close connection to peer: %s", err)
}
}
r.storeMu.Unlock()
}
func handShake(conn net.Conn) (*Peer, error) {
buf := make([]byte, 1500)
n, err := conn.Read(buf)
+11 -1
View File
@@ -24,7 +24,6 @@ func (s *Store) AddPeer(peer *Peer) {
func (s *Store) DeletePeer(peer *Peer) {
s.peersLock.Lock()
defer s.peersLock.Unlock()
delete(s.peers, peer.String())
}
@@ -35,3 +34,14 @@ func (s *Store) Peer(id string) (*Peer, bool) {
p, ok := s.peers[id]
return p, ok
}
func (s *Store) Peers() []*Peer {
s.peersLock.RLock()
defer s.peersLock.RUnlock()
peers := make([]*Peer, 0, len(s.peers))
for _, p := range s.peers {
peers = append(peers, p)
}
return peers
}