mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-03 20:19:07 +02:00
Remove channel binding logic
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user