Remove channel binding logic

This commit is contained in:
Zoltán Papp
2024-05-23 13:24:02 +02:00
parent 0a05f8b4d4
commit 36b2cd16cc
11 changed files with 229 additions and 374 deletions
+21 -15
View File
@@ -10,15 +10,15 @@ import (
)
type Listener struct {
address string
address string
conns map[string]*UDPConn
onAcceptFn func(conn net.Conn)
conns map[string]*UDPConn
wg sync.WaitGroup
quit chan struct{}
lock sync.Mutex
listener *net.UDPConn
wg sync.WaitGroup
quit chan struct{}
lock sync.Mutex
}
func NewListener(address string) listener.Listener {
@@ -34,17 +34,20 @@ func (l *Listener) Listen(onAcceptFn func(conn net.Conn)) error {
l.onAcceptFn = onAcceptFn
l.quit = make(chan struct{})
addr := &net.UDPAddr{
Port: 1234,
IP: net.ParseIP("0.0.0.0"),
}
li, err := net.ListenUDP("udp", addr)
addr, err := net.ResolveUDPAddr("udp", l.address)
if err != nil {
log.Errorf("%s", err)
log.Errorf("invalid listen address '%s': %s", l.address, err)
l.lock.Unlock()
return err
}
log.Debugf("udp server is listening on address: %s", l.address)
li, err := net.ListenUDP("udp", addr)
if err != nil {
log.Fatalf("%s", err)
l.lock.Unlock()
return err
}
log.Debugf("udp server is listening on address: %s", addr.String())
l.listener = li
l.wg.Add(1)
go l.readLoop()
@@ -54,14 +57,18 @@ func (l *Listener) Listen(onAcceptFn func(conn net.Conn)) error {
return nil
}
// Close todo: prevent multiple call (do not close two times the channel)
func (l *Listener) Close() error {
l.lock.Lock()
defer l.lock.Unlock()
if l.listener == nil {
return nil
}
close(l.quit)
err := l.listener.Close()
l.wg.Wait()
l.listener = nil
return err
}
@@ -91,6 +98,5 @@ func (l *Listener) readLoop() {
l.conns[addr.String()] = pConn
go l.onAcceptFn(pConn)
pConn.onNewMsg(buf[:n])
}
}
+15 -94
View File
@@ -1,113 +1,34 @@
package server
import (
"fmt"
"net"
"sync"
log "github.com/sirupsen/logrus"
)
type Participant struct {
ChannelID uint16
ChannelIDForeign uint16
ConnForeign net.Conn
Peer *Peer
}
"github.com/netbirdio/netbird/relay/messages"
)
type Peer struct {
Log *log.Entry
id string
idS string
idB []byte
conn net.Conn
pendingParticipantByChannelID map[uint16]*Participant
participantByID map[uint16]*Participant // used for package transfer
participantByPeerID map[string]*Participant // used for channel linking
lastId uint16
lastIdLock sync.Mutex
}
func NewPeer(id string, conn net.Conn) *Peer {
func NewPeer(id []byte, conn net.Conn) *Peer {
log.Debugf("new peer: %v", id)
stringID := messages.HashIDToString(id)
return &Peer{
Log: log.WithField("peer_id", id),
id: id,
conn: conn,
pendingParticipantByChannelID: make(map[uint16]*Participant),
participantByID: make(map[uint16]*Participant),
participantByPeerID: make(map[string]*Participant),
Log: log.WithField("peer_id", stringID),
idB: id,
idS: stringID,
conn: conn,
}
}
func (p *Peer) BindChannel(remotePeerId string) uint16 {
ch, ok := p.participantByPeerID[remotePeerId]
if ok {
return ch.ChannelID
}
channelID := p.newChannelID()
channel := &Participant{
ChannelID: channelID,
}
p.pendingParticipantByChannelID[channelID] = channel
p.participantByPeerID[remotePeerId] = channel
return channelID
func (p *Peer) ID() []byte {
return p.idB
}
func (p *Peer) UnBindChannel(remotePeerId string) {
pa, ok := p.participantByPeerID[remotePeerId]
if !ok {
return
}
p.Log.Debugf("unbind channel with '%s': %d", remotePeerId, pa.ChannelID)
p.pendingParticipantByChannelID[pa.ChannelID] = pa
delete(p.participantByID, pa.ChannelID)
}
func (p *Peer) AddParticipant(peer *Peer, remoteChannelID uint16) (uint16, bool) {
participant, ok := p.participantByPeerID[peer.ID()]
if !ok {
return 0, false
}
participant.ChannelIDForeign = remoteChannelID
participant.ConnForeign = peer.conn
participant.Peer = peer
delete(p.pendingParticipantByChannelID, participant.ChannelID)
p.participantByID[participant.ChannelID] = participant
return participant.ChannelID, true
}
func (p *Peer) DeleteParticipants() {
for _, participant := range p.participantByID {
participant.Peer.UnBindChannel(p.id)
}
}
func (p *Peer) ConnByChannelID(dstID uint16) (uint16, net.Conn, error) {
ch, ok := p.participantByID[dstID]
if !ok {
return 0, nil, fmt.Errorf("destination channel not found")
}
return ch.ChannelIDForeign, ch.ConnForeign, nil
}
func (p *Peer) ID() string {
return p.id
}
func (p *Peer) newChannelID() uint16 {
p.lastIdLock.Lock()
defer p.lastIdLock.Unlock()
for {
p.lastId++
if _, ok := p.pendingParticipantByChannelID[p.lastId]; ok {
continue
}
if _, ok := p.participantByID[p.lastId]; ok {
continue
}
return p.lastId
}
func (p *Peer) String() string {
return p.idS
}
+13 -33
View File
@@ -16,7 +16,6 @@ import (
// todo:
// authentication: provide JWT token via RPC call. The MGM server can forward the token to the agents.
// connection timeout handling
// implement HA (High Availability) mode
type Server struct {
store *Store
@@ -75,54 +74,35 @@ func (r *Server) accept(conn net.Conn) {
return
}
switch msgType {
case messages.MsgTypeBindNewChannel:
dstPeerId, err := messages.UnmarshalBindNewChannel(buf[:n])
if err != nil {
log.Errorf("failed to unmarshal bind new channel message: %s", err)
continue
}
channelID := r.store.Link(peer, dstPeerId)
msg := messages.MarshalBindResponseMsg(channelID, dstPeerId)
_, err = conn.Write(msg)
if err != nil {
peer.Log.Errorf("failed to response to bind request: %s", err)
continue
}
peer.Log.Debugf("bind new channel with '%s', channelID: %d", dstPeerId, channelID)
case messages.MsgTypeTransport:
msg := buf[:n]
channelId, err := messages.UnmarshalTransportID(msg)
peerID, err := messages.UnmarshalTransportID(msg)
if err != nil {
peer.Log.Errorf("failed to unmarshal transport message: %s", err)
continue
}
go func() {
foreignChannelID, remoteConn, err := peer.ConnByChannelID(channelId)
if err != nil {
peer.Log.Errorf("failed to transport message from peer '%s' to '%d': %s", peer.ID(), channelId, err)
stringPeerID := messages.HashIDToString(peerID)
dp, ok := r.store.Peer(stringPeerID)
if !ok {
peer.Log.Errorf("peer not found: %s", stringPeerID)
return
}
err = transportTo(remoteConn, foreignChannelID, msg)
err := messages.UpdateTransportMsg(msg, peer.ID())
if err != nil {
peer.Log.Errorf("failed to transport message from peer '%s' to '%d': %s", peer.ID(), channelId, err)
peer.Log.Errorf("failed to update transport message: %s", err)
return
}
_, err = dp.conn.Write(msg)
if err != nil {
peer.Log.Errorf("failed to write transport message to: %s", dp.String())
}
return
}()
}
}
}
func transportTo(conn net.Conn, channelID uint16, msg []byte) error {
err := messages.UpdateTransportMsg(msg, channelID)
if err != nil {
return err
}
_, err = conn.Write(msg)
return err
}
func handShake(conn net.Conn) (*Peer, error) {
buf := make([]byte, 1500)
n, err := conn.Read(buf)
+12 -23
View File
@@ -5,8 +5,8 @@ import (
)
type Store struct {
peers map[string]*Peer // Key is the id (public key or sha-256) of the peer
peersLock sync.Mutex
peers map[string]*Peer // consider to use [32]byte as key. The Peer(id string) would be faster
peersLock sync.RWMutex
}
func NewStore() *Store {
@@ -18,31 +18,20 @@ func NewStore() *Store {
func (s *Store) AddPeer(peer *Peer) {
s.peersLock.Lock()
defer s.peersLock.Unlock()
s.peers[peer.ID()] = peer
}
func (s *Store) Link(peer *Peer, peerForeignID string) uint16 {
s.peersLock.Lock()
defer s.peersLock.Unlock()
channelId := peer.BindChannel(peerForeignID)
dstPeer, ok := s.peers[peerForeignID]
if !ok {
return channelId
}
foreignChannelID, ok := dstPeer.AddParticipant(peer, channelId)
if !ok {
return channelId
}
peer.AddParticipant(dstPeer, foreignChannelID)
return channelId
s.peers[peer.String()] = peer
}
func (s *Store) DeletePeer(peer *Peer) {
s.peersLock.Lock()
defer s.peersLock.Unlock()
delete(s.peers, peer.ID())
peer.DeleteParticipants()
delete(s.peers, peer.String())
}
func (s *Store) Peer(id string) (*Peer, bool) {
s.peersLock.RLock()
defer s.peersLock.RUnlock()
p, ok := s.peers[id]
return p, ok
}