186 lines
3.6 KiB
Go
186 lines
3.6 KiB
Go
// Package discovery provides optional mDNS-based peer auto-discovery.
|
|
//
|
|
// Each instance advertises a `_clip-sync._tcp` service and periodically browses
|
|
// for other instances, adding them to (and removing them from) the broadcaster.
|
|
package discovery
|
|
|
|
import (
|
|
"log"
|
|
"net"
|
|
"os"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/hashicorp/mdns"
|
|
|
|
"git.dracodev.net/Projets/clip-sync/internal/config"
|
|
"git.dracodev.net/Projets/clip-sync/internal/peer"
|
|
)
|
|
|
|
const (
|
|
serviceName = "_clip-sync._tcp"
|
|
domain = "local"
|
|
)
|
|
|
|
// Discovery manages mDNS advertisement and browsing.
|
|
type Discovery interface {
|
|
Start() error
|
|
Stop()
|
|
}
|
|
|
|
type mdnsDiscovery struct {
|
|
origin string
|
|
port int
|
|
broadcaster *peer.Broadcaster
|
|
|
|
mu sync.Mutex
|
|
server *mdns.Server
|
|
stopCh chan struct{}
|
|
doneCh chan struct{}
|
|
peers map[string]string // addr -> host
|
|
}
|
|
|
|
// New creates an mDNS discovery component that updates the given broadcaster.
|
|
func New(origin string, port int, b *peer.Broadcaster) Discovery {
|
|
return &mdnsDiscovery{
|
|
origin: origin,
|
|
port: port,
|
|
broadcaster: b,
|
|
peers: make(map[string]string),
|
|
}
|
|
}
|
|
|
|
func (d *mdnsDiscovery) Start() error {
|
|
host, _ := os.Hostname()
|
|
instance := "clip-sync-" + sanitize(host)
|
|
|
|
// Advertise this instance. Zero-value domain/host/ips are inferred by the
|
|
// library from the operating system.
|
|
service, err := mdns.NewMDNSService(instance, serviceName, "", "", d.port, nil, nil)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
server, err := mdns.NewServer(&mdns.Config{Zone: service})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
d.mu.Lock()
|
|
d.server = server
|
|
d.stopCh = make(chan struct{})
|
|
d.doneCh = make(chan struct{})
|
|
d.mu.Unlock()
|
|
|
|
go d.browseLoop()
|
|
log.Printf("discovery: mDNS advertised as %q, browsing for peers", instance)
|
|
return nil
|
|
}
|
|
|
|
func (d *mdnsDiscovery) Stop() {
|
|
d.mu.Lock()
|
|
stopCh := d.stopCh
|
|
server := d.server
|
|
d.mu.Unlock()
|
|
|
|
if stopCh != nil {
|
|
close(stopCh)
|
|
}
|
|
if server != nil {
|
|
_ = server.Shutdown()
|
|
}
|
|
if d.doneCh != nil {
|
|
<-d.doneCh
|
|
}
|
|
}
|
|
|
|
func (d *mdnsDiscovery) browseLoop() {
|
|
defer close(d.doneCh)
|
|
|
|
ticker := time.NewTicker(30 * time.Second)
|
|
defer ticker.Stop()
|
|
|
|
d.browse()
|
|
for {
|
|
select {
|
|
case <-d.stopCh:
|
|
return
|
|
case <-ticker.C:
|
|
d.browse()
|
|
}
|
|
}
|
|
}
|
|
|
|
func (d *mdnsDiscovery) browse() {
|
|
entries := make(chan *mdns.ServiceEntry, 16)
|
|
params := &mdns.QueryParam{
|
|
Service: serviceName,
|
|
Domain: domain,
|
|
Timeout: 5 * time.Second,
|
|
Entries: entries,
|
|
}
|
|
|
|
seen := make(map[string]string)
|
|
go func() {
|
|
_ = mdns.Query(params)
|
|
close(entries)
|
|
}()
|
|
|
|
for entry := range entries {
|
|
host := strings.TrimSuffix(entry.Host, ".local.")
|
|
host = strings.TrimSuffix(host, ".")
|
|
|
|
// Skip ourselves (and our sanitized advertisement form).
|
|
if host == d.origin || host == sanitize(d.origin) {
|
|
continue
|
|
}
|
|
|
|
ip := entry.AddrV4
|
|
if ip == nil {
|
|
ip = entry.AddrV6
|
|
}
|
|
if ip == nil {
|
|
continue
|
|
}
|
|
|
|
addr := net.JoinHostPort(ip.String(), itoa(entry.Port))
|
|
seen[addr] = host
|
|
d.broadcaster.AddPeer(config.PeerConfig{Name: host, Addr: addr})
|
|
}
|
|
|
|
// Remove peers that disappeared since the last browse.
|
|
d.mu.Lock()
|
|
for addr := range d.peers {
|
|
if _, ok := seen[addr]; !ok {
|
|
d.broadcaster.RemovePeer(addr)
|
|
}
|
|
}
|
|
d.peers = seen
|
|
d.mu.Unlock()
|
|
}
|
|
|
|
func sanitize(s string) string {
|
|
return strings.Map(func(r rune) rune {
|
|
switch {
|
|
case r >= 'a' && r <= 'z', r >= 'A' && r <= 'Z', r >= '0' && r <= '9', r == '-':
|
|
return r
|
|
default:
|
|
return '-'
|
|
}
|
|
}, s)
|
|
}
|
|
|
|
func itoa(n int) string {
|
|
if n == 0 {
|
|
return "0"
|
|
}
|
|
var b [20]byte
|
|
i := len(b)
|
|
for n > 0 {
|
|
i--
|
|
b[i] = byte('0' + n%10)
|
|
n /= 10
|
|
}
|
|
return string(b[i:])
|
|
}
|