mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: use tailnet v2 API for coordination (#11638)
This one is huge, and I'm sorry. The problem is that once I change `tailnet.Conn` to start doing v2 behavior, I kind of have to change it everywhere, including in CoderSDK (CLI), the agent, wsproxy, and ServerTailnet. There is still a bit more cleanup to do, and I need to add code so that when we lose connection to the Coordinator, we mark all peers as LOST, but that will be in a separate PR since this is big enough!
This commit is contained in:
@@ -490,6 +490,18 @@ func (c *configMaps) protoNodeToTailcfg(p *proto.Node) (*tailcfg.Node, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// nodeAddresses returns the addresses for the peer with the given publicKey, if known.
|
||||
func (c *configMaps) nodeAddresses(publicKey key.NodePublic) ([]netip.Prefix, bool) {
|
||||
c.L.Lock()
|
||||
defer c.L.Unlock()
|
||||
for _, lc := range c.peers {
|
||||
if lc.node.Key == publicKey {
|
||||
return lc.node.Addresses, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
type peerLifecycle struct {
|
||||
peerID uuid.UUID
|
||||
node *tailcfg.Node
|
||||
|
||||
+62
-457
@@ -3,48 +3,40 @@ package tailnet
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"os"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/cenkalti/backoff/v4"
|
||||
"github.com/google/uuid"
|
||||
"go4.org/netipx"
|
||||
"golang.org/x/xerrors"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
|
||||
"tailscale.com/envknob"
|
||||
"tailscale.com/ipn/ipnstate"
|
||||
"tailscale.com/net/connstats"
|
||||
"tailscale.com/net/dns"
|
||||
"tailscale.com/net/netmon"
|
||||
"tailscale.com/net/netns"
|
||||
"tailscale.com/net/tsdial"
|
||||
"tailscale.com/net/tstun"
|
||||
"tailscale.com/tailcfg"
|
||||
"tailscale.com/tsd"
|
||||
"tailscale.com/types/ipproto"
|
||||
"tailscale.com/types/key"
|
||||
tslogger "tailscale.com/types/logger"
|
||||
"tailscale.com/types/netlogtype"
|
||||
"tailscale.com/types/netmap"
|
||||
"tailscale.com/wgengine"
|
||||
"tailscale.com/wgengine/filter"
|
||||
"tailscale.com/wgengine/magicsock"
|
||||
"tailscale.com/wgengine/netstack"
|
||||
"tailscale.com/wgengine/router"
|
||||
"tailscale.com/wgengine/wgcfg/nmcfg"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/cryptorand"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
)
|
||||
|
||||
var ErrConnClosed = xerrors.New("connection closed")
|
||||
@@ -128,42 +120,6 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
}
|
||||
|
||||
nodePrivateKey := key.NewNode()
|
||||
nodePublicKey := nodePrivateKey.Public()
|
||||
|
||||
netMap := &netmap.NetworkMap{
|
||||
DERPMap: options.DERPMap,
|
||||
NodeKey: nodePublicKey,
|
||||
PrivateKey: nodePrivateKey,
|
||||
Addresses: options.Addresses,
|
||||
PacketFilter: []filter.Match{{
|
||||
// Allow any protocol!
|
||||
IPProto: []ipproto.Proto{ipproto.TCP, ipproto.UDP, ipproto.ICMPv4, ipproto.ICMPv6, ipproto.SCTP},
|
||||
// Allow traffic sourced from anywhere.
|
||||
Srcs: []netip.Prefix{
|
||||
netip.PrefixFrom(netip.AddrFrom4([4]byte{}), 0),
|
||||
netip.PrefixFrom(netip.AddrFrom16([16]byte{}), 0),
|
||||
},
|
||||
// Allow traffic to route anywhere.
|
||||
Dsts: []filter.NetPortRange{
|
||||
{
|
||||
Net: netip.PrefixFrom(netip.AddrFrom4([4]byte{}), 0),
|
||||
Ports: filter.PortRange{
|
||||
First: 0,
|
||||
Last: 65535,
|
||||
},
|
||||
},
|
||||
{
|
||||
Net: netip.PrefixFrom(netip.AddrFrom16([16]byte{}), 0),
|
||||
Ports: filter.PortRange{
|
||||
First: 0,
|
||||
Last: 65535,
|
||||
},
|
||||
},
|
||||
},
|
||||
Caps: []filter.CapMatch{},
|
||||
}},
|
||||
}
|
||||
|
||||
var nodeID tailcfg.NodeID
|
||||
|
||||
// If we're provided with a UUID, use it to populate our node ID.
|
||||
@@ -177,14 +133,6 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
nodeID = tailcfg.NodeID(uid)
|
||||
}
|
||||
|
||||
// This is used by functions below to identify the node via key
|
||||
netMap.SelfNode = &tailcfg.Node{
|
||||
ID: nodeID,
|
||||
Key: nodePublicKey,
|
||||
Addresses: options.Addresses,
|
||||
AllowedIPs: options.Addresses,
|
||||
}
|
||||
|
||||
wireguardMonitor, err := netmon.New(Logger(options.Logger.Named("net.wgmonitor")))
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("create wireguard link monitor: %w", err)
|
||||
@@ -243,7 +191,6 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("set node private key: %w", err)
|
||||
}
|
||||
netMap.SelfNode.DiscoKey = magicConn.DiscoPublicKey()
|
||||
|
||||
netStack, err := netstack.Create(
|
||||
Logger(options.Logger.Named("net.netstack")),
|
||||
@@ -262,44 +209,46 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
}
|
||||
netStack.ProcessLocalIPs = true
|
||||
wireguardEngine = wgengine.NewWatchdog(wireguardEngine)
|
||||
wireguardEngine.SetDERPMap(options.DERPMap)
|
||||
netMapCopy := *netMap
|
||||
options.Logger.Debug(context.Background(), "updating network map")
|
||||
wireguardEngine.SetNetworkMap(&netMapCopy)
|
||||
|
||||
localIPSet := netipx.IPSetBuilder{}
|
||||
for _, addr := range netMap.Addresses {
|
||||
localIPSet.AddPrefix(addr)
|
||||
}
|
||||
localIPs, _ := localIPSet.IPSet()
|
||||
logIPSet := netipx.IPSetBuilder{}
|
||||
logIPs, _ := logIPSet.IPSet()
|
||||
wireguardEngine.SetFilter(filter.New(
|
||||
netMap.PacketFilter,
|
||||
localIPs,
|
||||
logIPs,
|
||||
cfgMaps := newConfigMaps(
|
||||
options.Logger,
|
||||
wireguardEngine,
|
||||
nodeID,
|
||||
nodePrivateKey,
|
||||
magicConn.DiscoPublicKey(),
|
||||
)
|
||||
cfgMaps.setAddresses(options.Addresses)
|
||||
cfgMaps.setDERPMap(DERPMapToProto(options.DERPMap))
|
||||
cfgMaps.setBlockEndpoints(options.BlockEndpoints)
|
||||
|
||||
nodeUp := newNodeUpdater(
|
||||
options.Logger,
|
||||
nil,
|
||||
Logger(options.Logger.Named("net.packet-filter")),
|
||||
))
|
||||
nodeID,
|
||||
nodePrivateKey.Public(),
|
||||
magicConn.DiscoPublicKey(),
|
||||
)
|
||||
nodeUp.setAddresses(options.Addresses)
|
||||
nodeUp.setBlockEndpoints(options.BlockEndpoints)
|
||||
wireguardEngine.SetStatusCallback(nodeUp.setStatus)
|
||||
wireguardEngine.SetNetInfoCallback(nodeUp.setNetInfo)
|
||||
magicConn.SetDERPForcedWebsocketCallback(nodeUp.setDERPForcedWebsocket)
|
||||
|
||||
server := &Conn{
|
||||
blockEndpoints: options.BlockEndpoints,
|
||||
derpForceWebSockets: options.DERPForceWebSockets,
|
||||
closed: make(chan struct{}),
|
||||
logger: options.Logger,
|
||||
magicConn: magicConn,
|
||||
dialer: dialer,
|
||||
listeners: map[listenKey]*listener{},
|
||||
peerMap: map[tailcfg.NodeID]*tailcfg.Node{},
|
||||
lastDERPForcedWebSockets: map[int]string{},
|
||||
tunDevice: sys.Tun.Get(),
|
||||
netMap: netMap,
|
||||
netStack: netStack,
|
||||
wireguardMonitor: wireguardMonitor,
|
||||
closed: make(chan struct{}),
|
||||
logger: options.Logger,
|
||||
magicConn: magicConn,
|
||||
dialer: dialer,
|
||||
listeners: map[listenKey]*listener{},
|
||||
tunDevice: sys.Tun.Get(),
|
||||
netStack: netStack,
|
||||
wireguardMonitor: wireguardMonitor,
|
||||
wireguardRouter: &router.Config{
|
||||
LocalAddrs: netMap.Addresses,
|
||||
LocalAddrs: options.Addresses,
|
||||
},
|
||||
wireguardEngine: wireguardEngine,
|
||||
configMaps: cfgMaps,
|
||||
nodeUpdater: nodeUp,
|
||||
}
|
||||
defer func() {
|
||||
if err != nil {
|
||||
@@ -307,52 +256,6 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
}
|
||||
}()
|
||||
|
||||
wireguardEngine.SetStatusCallback(func(s *wgengine.Status, err error) {
|
||||
server.logger.Debug(context.Background(), "wireguard status", slog.F("status", s), slog.Error(err))
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
server.lastMutex.Lock()
|
||||
if s.AsOf.Before(server.lastStatus) {
|
||||
// Don't process outdated status!
|
||||
server.lastMutex.Unlock()
|
||||
return
|
||||
}
|
||||
server.lastStatus = s.AsOf
|
||||
if endpointsEqual(s.LocalAddrs, server.lastEndpoints) {
|
||||
// No need to update the node if nothing changed!
|
||||
server.lastMutex.Unlock()
|
||||
return
|
||||
}
|
||||
server.lastEndpoints = append([]tailcfg.Endpoint{}, s.LocalAddrs...)
|
||||
server.lastMutex.Unlock()
|
||||
server.sendNode()
|
||||
})
|
||||
|
||||
wireguardEngine.SetNetInfoCallback(func(ni *tailcfg.NetInfo) {
|
||||
server.logger.Debug(context.Background(), "netinfo callback", slog.F("netinfo", ni))
|
||||
server.lastMutex.Lock()
|
||||
if reflect.DeepEqual(server.lastNetInfo, ni) {
|
||||
server.lastMutex.Unlock()
|
||||
return
|
||||
}
|
||||
server.lastNetInfo = ni.Clone()
|
||||
server.lastMutex.Unlock()
|
||||
server.sendNode()
|
||||
})
|
||||
|
||||
magicConn.SetDERPForcedWebsocketCallback(func(region int, reason string) {
|
||||
server.logger.Debug(context.Background(), "derp forced websocket", slog.F("region", region), slog.F("reason", reason))
|
||||
server.lastMutex.Lock()
|
||||
if server.lastDERPForcedWebSockets[region] == reason {
|
||||
server.lastMutex.Unlock()
|
||||
return
|
||||
}
|
||||
server.lastDERPForcedWebSockets[region] = reason
|
||||
server.lastMutex.Unlock()
|
||||
server.sendNode()
|
||||
})
|
||||
|
||||
netStack.GetTCPHandlerForFlow = server.forwardTCP
|
||||
|
||||
err = netStack.Start(nil)
|
||||
@@ -389,16 +292,14 @@ func IPFromUUID(uid uuid.UUID) netip.Addr {
|
||||
|
||||
// Conn is an actively listening Wireguard connection.
|
||||
type Conn struct {
|
||||
mutex sync.Mutex
|
||||
closed chan struct{}
|
||||
logger slog.Logger
|
||||
blockEndpoints bool
|
||||
derpForceWebSockets bool
|
||||
mutex sync.Mutex
|
||||
closed chan struct{}
|
||||
logger slog.Logger
|
||||
|
||||
dialer *tsdial.Dialer
|
||||
tunDevice *tstun.Wrapper
|
||||
peerMap map[tailcfg.NodeID]*tailcfg.Node
|
||||
netMap *netmap.NetworkMap
|
||||
configMaps *configMaps
|
||||
nodeUpdater *nodeUpdater
|
||||
netStack *netstack.Impl
|
||||
magicConn *magicsock.Conn
|
||||
wireguardMonitor *netmon.Monitor
|
||||
@@ -406,17 +307,6 @@ type Conn struct {
|
||||
wireguardEngine wgengine.Engine
|
||||
listeners map[listenKey]*listener
|
||||
|
||||
lastMutex sync.Mutex
|
||||
nodeSending bool
|
||||
nodeChanged bool
|
||||
// It's only possible to store these values via status functions,
|
||||
// so the values must be stored for retrieval later on.
|
||||
lastStatus time.Time
|
||||
lastEndpoints []tailcfg.Endpoint
|
||||
lastDERPForcedWebSockets map[int]string
|
||||
lastNetInfo *tailcfg.NetInfo
|
||||
nodeCallback func(node *Node)
|
||||
|
||||
trafficStats *connstats.Statistics
|
||||
}
|
||||
|
||||
@@ -425,57 +315,30 @@ func (c *Conn) MagicsockSetDebugLoggingEnabled(enabled bool) {
|
||||
}
|
||||
|
||||
func (c *Conn) SetAddresses(ips []netip.Prefix) error {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
c.netMap.Addresses = ips
|
||||
|
||||
netMapCopy := *c.netMap
|
||||
c.logger.Debug(context.Background(), "updating network map")
|
||||
c.wireguardEngine.SetNetworkMap(&netMapCopy)
|
||||
err := c.reconfig()
|
||||
if err != nil {
|
||||
return xerrors.Errorf("reconfig: %w", err)
|
||||
}
|
||||
|
||||
c.configMaps.setAddresses(ips)
|
||||
c.nodeUpdater.setAddresses(ips)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) Addresses() []netip.Prefix {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
return c.netMap.Addresses
|
||||
}
|
||||
|
||||
func (c *Conn) SetNodeCallback(callback func(node *Node)) {
|
||||
c.lastMutex.Lock()
|
||||
c.nodeCallback = callback
|
||||
c.lastMutex.Unlock()
|
||||
c.sendNode()
|
||||
c.nodeUpdater.setCallback(callback)
|
||||
}
|
||||
|
||||
// SetDERPMap updates the DERPMap of a connection.
|
||||
func (c *Conn) SetDERPMap(derpMap *tailcfg.DERPMap) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
c.logger.Debug(context.Background(), "updating derp map", slog.F("derp_map", derpMap))
|
||||
c.wireguardEngine.SetDERPMap(derpMap)
|
||||
c.netMap.DERPMap = derpMap
|
||||
netMapCopy := *c.netMap
|
||||
c.logger.Debug(context.Background(), "updating network map")
|
||||
c.wireguardEngine.SetNetworkMap(&netMapCopy)
|
||||
c.configMaps.setDERPMap(DERPMapToProto(derpMap))
|
||||
}
|
||||
|
||||
func (c *Conn) SetDERPForceWebSockets(v bool) {
|
||||
c.logger.Info(context.Background(), "setting DERP Force Websockets", slog.F("force_derp_websockets", v))
|
||||
c.magicConn.SetDERPForceWebsockets(v)
|
||||
}
|
||||
|
||||
// SetBlockEndpoints sets whether or not to block P2P endpoints. This setting
|
||||
// SetBlockEndpoints sets whether to block P2P endpoints. This setting
|
||||
// will only apply to new peers.
|
||||
func (c *Conn) SetBlockEndpoints(blockEndpoints bool) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
c.blockEndpoints = blockEndpoints
|
||||
c.configMaps.setBlockEndpoints(blockEndpoints)
|
||||
c.nodeUpdater.setBlockEndpoints(blockEndpoints)
|
||||
}
|
||||
|
||||
// SetDERPRegionDialer updates the dialer to use for connecting to DERP regions.
|
||||
@@ -483,186 +346,24 @@ func (c *Conn) SetDERPRegionDialer(dialer func(ctx context.Context, region *tail
|
||||
c.magicConn.SetDERPRegionDialer(dialer)
|
||||
}
|
||||
|
||||
// UpdateNodes connects with a set of peers. This can be constantly updated,
|
||||
// and peers will continually be reconnected as necessary. If replacePeers is
|
||||
// true, all peers will be removed before adding the new ones.
|
||||
//
|
||||
//nolint:revive // Complains about replacePeers.
|
||||
func (c *Conn) UpdateNodes(nodes []*Node, replacePeers bool) error {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
// UpdatePeers connects with a set of peers. This can be constantly updated,
|
||||
// and peers will continually be reconnected as necessary.
|
||||
func (c *Conn) UpdatePeers(updates []*proto.CoordinateResponse_PeerUpdate) error {
|
||||
if c.isClosed() {
|
||||
return ErrConnClosed
|
||||
}
|
||||
|
||||
status := c.Status()
|
||||
if replacePeers {
|
||||
c.netMap.Peers = []*tailcfg.Node{}
|
||||
c.peerMap = map[tailcfg.NodeID]*tailcfg.Node{}
|
||||
}
|
||||
for _, peer := range c.netMap.Peers {
|
||||
peerStatus, ok := status.Peer[peer.Key]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
// If this peer was added in the last 5 minutes, assume it
|
||||
// could still be active.
|
||||
if time.Since(peer.Created) < 5*time.Minute {
|
||||
continue
|
||||
}
|
||||
// We double-check that it's safe to remove by ensuring no
|
||||
// handshake has been sent in the past 5 minutes as well. Connections that
|
||||
// are actively exchanging IP traffic will handshake every 2 minutes.
|
||||
if time.Since(peerStatus.LastHandshake) < 5*time.Minute {
|
||||
continue
|
||||
}
|
||||
|
||||
c.logger.Debug(context.Background(), "removing peer, last handshake >5m ago",
|
||||
slog.F("peer", peer.Key), slog.F("last_handshake", peerStatus.LastHandshake),
|
||||
)
|
||||
delete(c.peerMap, peer.ID)
|
||||
}
|
||||
|
||||
for _, node := range nodes {
|
||||
// If no preferred DERP is provided, we can't reach the node.
|
||||
if node.PreferredDERP == 0 {
|
||||
c.logger.Debug(context.Background(), "no preferred DERP, skipping node", slog.F("node", node))
|
||||
continue
|
||||
}
|
||||
c.logger.Debug(context.Background(), "adding node", slog.F("node", node))
|
||||
|
||||
peerStatus, ok := status.Peer[node.Key]
|
||||
peerNode := &tailcfg.Node{
|
||||
ID: node.ID,
|
||||
Created: time.Now(),
|
||||
Key: node.Key,
|
||||
DiscoKey: node.DiscoKey,
|
||||
Addresses: node.Addresses,
|
||||
AllowedIPs: node.AllowedIPs,
|
||||
Endpoints: node.Endpoints,
|
||||
DERP: fmt.Sprintf("%s:%d", tailcfg.DerpMagicIP, node.PreferredDERP),
|
||||
Hostinfo: (&tailcfg.Hostinfo{}).View(),
|
||||
// Starting KeepAlive messages at the initialization of a connection
|
||||
// causes a race condition. If we handshake before the peer has our
|
||||
// node, we'll have wait for 5 seconds before trying again. Ideally,
|
||||
// the first handshake starts when the user first initiates a
|
||||
// connection to the peer. After a successful connection we enable
|
||||
// keep alives to persist the connection and keep it from becoming
|
||||
// idle. SSH connections don't send send packets while idle, so we
|
||||
// use keep alives to avoid random hangs while we set up the
|
||||
// connection again after inactivity.
|
||||
KeepAlive: ok && peerStatus.Active,
|
||||
}
|
||||
if c.blockEndpoints {
|
||||
peerNode.Endpoints = nil
|
||||
}
|
||||
c.peerMap[node.ID] = peerNode
|
||||
}
|
||||
|
||||
c.netMap.Peers = make([]*tailcfg.Node, 0, len(c.peerMap))
|
||||
for _, peer := range c.peerMap {
|
||||
c.netMap.Peers = append(c.netMap.Peers, peer.Clone())
|
||||
}
|
||||
|
||||
netMapCopy := *c.netMap
|
||||
c.logger.Debug(context.Background(), "updating network map")
|
||||
c.wireguardEngine.SetNetworkMap(&netMapCopy)
|
||||
err := c.reconfig()
|
||||
if err != nil {
|
||||
return xerrors.Errorf("reconfig: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// PeerSelector is used to select a peer from within a Tailnet.
|
||||
type PeerSelector struct {
|
||||
ID tailcfg.NodeID
|
||||
IP netip.Prefix
|
||||
}
|
||||
|
||||
func (c *Conn) RemovePeer(selector PeerSelector) (deleted bool, err error) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if c.isClosed() {
|
||||
return false, ErrConnClosed
|
||||
}
|
||||
|
||||
deleted = false
|
||||
for _, peer := range c.peerMap {
|
||||
if peer.ID == selector.ID {
|
||||
delete(c.peerMap, peer.ID)
|
||||
deleted = true
|
||||
break
|
||||
}
|
||||
|
||||
for _, peerIP := range peer.Addresses {
|
||||
if peerIP.Bits() == selector.IP.Bits() && peerIP.Addr().Compare(selector.IP.Addr()) == 0 {
|
||||
delete(c.peerMap, peer.ID)
|
||||
deleted = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if !deleted {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
c.netMap.Peers = make([]*tailcfg.Node, 0, len(c.peerMap))
|
||||
for _, peer := range c.peerMap {
|
||||
c.netMap.Peers = append(c.netMap.Peers, peer.Clone())
|
||||
}
|
||||
|
||||
netMapCopy := *c.netMap
|
||||
c.logger.Debug(context.Background(), "updating network map")
|
||||
c.wireguardEngine.SetNetworkMap(&netMapCopy)
|
||||
err = c.reconfig()
|
||||
if err != nil {
|
||||
return false, xerrors.Errorf("reconfig: %w", err)
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (c *Conn) reconfig() error {
|
||||
cfg, err := nmcfg.WGCfg(c.netMap, Logger(c.logger.Named("net.wgconfig")), netmap.AllowSingleHosts, "")
|
||||
if err != nil {
|
||||
return xerrors.Errorf("update wireguard config: %w", err)
|
||||
}
|
||||
|
||||
err = c.wireguardEngine.Reconfig(cfg, c.wireguardRouter, &dns.Config{}, &tailcfg.Debug{})
|
||||
if err != nil {
|
||||
if c.isClosed() {
|
||||
return nil
|
||||
}
|
||||
if errors.Is(err, wgengine.ErrNoChanges) {
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("reconfig: %w", err)
|
||||
}
|
||||
|
||||
c.configMaps.updatePeers(updates)
|
||||
return nil
|
||||
}
|
||||
|
||||
// NodeAddresses returns the addresses of a node from the NetworkMap.
|
||||
func (c *Conn) NodeAddresses(publicKey key.NodePublic) ([]netip.Prefix, bool) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
for _, node := range c.netMap.Peers {
|
||||
if node.Key == publicKey {
|
||||
return node.Addresses, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
return c.configMaps.nodeAddresses(publicKey)
|
||||
}
|
||||
|
||||
// Status returns the current ipnstate of a connection.
|
||||
func (c *Conn) Status() *ipnstate.Status {
|
||||
sb := &ipnstate.StatusBuilder{WantPeers: true}
|
||||
c.wireguardEngine.UpdateStatus(sb)
|
||||
return sb.Status()
|
||||
return c.configMaps.status()
|
||||
}
|
||||
|
||||
// Ping sends a ping to the Wireguard engine.
|
||||
@@ -689,16 +390,9 @@ func (c *Conn) Ping(ctx context.Context, ip netip.Addr) (time.Duration, bool, *i
|
||||
|
||||
// DERPMap returns the currently set DERP mapping.
|
||||
func (c *Conn) DERPMap() *tailcfg.DERPMap {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
return c.netMap.DERPMap
|
||||
}
|
||||
|
||||
// BlockEndpoints returns whether or not P2P is blocked.
|
||||
func (c *Conn) BlockEndpoints() bool {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
return c.blockEndpoints
|
||||
c.configMaps.L.Lock()
|
||||
defer c.configMaps.L.Unlock()
|
||||
return c.configMaps.derpMapLocked()
|
||||
}
|
||||
|
||||
// AwaitReachable pings the provided IP continually until the
|
||||
@@ -759,6 +453,9 @@ func (c *Conn) Closed() <-chan struct{} {
|
||||
|
||||
// Close shuts down the Wireguard connection.
|
||||
func (c *Conn) Close() error {
|
||||
c.logger.Info(context.Background(), "closing tailnet Conn")
|
||||
c.configMaps.close()
|
||||
c.nodeUpdater.close()
|
||||
c.mutex.Lock()
|
||||
select {
|
||||
case <-c.closed:
|
||||
@@ -808,91 +505,11 @@ func (c *Conn) isClosed() bool {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) sendNode() {
|
||||
c.lastMutex.Lock()
|
||||
defer c.lastMutex.Unlock()
|
||||
if c.nodeSending {
|
||||
c.nodeChanged = true
|
||||
return
|
||||
}
|
||||
node := c.selfNode()
|
||||
// Conn.UpdateNodes will skip any nodes that don't have the PreferredDERP
|
||||
// set to non-zero, since we cannot reach nodes without DERP for discovery.
|
||||
// Therefore, there is no point in sending the node without this, and we can
|
||||
// save ourselves from churn in the tailscale/wireguard layer.
|
||||
if node.PreferredDERP == 0 {
|
||||
c.logger.Debug(context.Background(), "skipped sending node; no PreferredDERP", slog.F("node", node))
|
||||
return
|
||||
}
|
||||
nodeCallback := c.nodeCallback
|
||||
if nodeCallback == nil {
|
||||
return
|
||||
}
|
||||
c.nodeSending = true
|
||||
go func() {
|
||||
c.logger.Debug(context.Background(), "sending node", slog.F("node", node))
|
||||
nodeCallback(node)
|
||||
c.lastMutex.Lock()
|
||||
c.nodeSending = false
|
||||
if c.nodeChanged {
|
||||
c.nodeChanged = false
|
||||
c.lastMutex.Unlock()
|
||||
c.sendNode()
|
||||
return
|
||||
}
|
||||
c.lastMutex.Unlock()
|
||||
}()
|
||||
}
|
||||
|
||||
// Node returns the last node that was sent to the node callback.
|
||||
func (c *Conn) Node() *Node {
|
||||
c.lastMutex.Lock()
|
||||
defer c.lastMutex.Unlock()
|
||||
return c.selfNode()
|
||||
}
|
||||
|
||||
func (c *Conn) selfNode() *Node {
|
||||
endpoints := make([]string, 0, len(c.lastEndpoints))
|
||||
for _, addr := range c.lastEndpoints {
|
||||
endpoints = append(endpoints, addr.Addr.String())
|
||||
}
|
||||
var preferredDERP int
|
||||
var derpLatency map[string]float64
|
||||
derpForcedWebsocket := make(map[int]string, 0)
|
||||
if c.lastNetInfo != nil {
|
||||
preferredDERP = c.lastNetInfo.PreferredDERP
|
||||
derpLatency = c.lastNetInfo.DERPLatency
|
||||
|
||||
if c.derpForceWebSockets {
|
||||
// We only need to store this for a single region, since this is
|
||||
// mostly used for debugging purposes and doesn't actually have a
|
||||
// code purpose.
|
||||
derpForcedWebsocket[preferredDERP] = "DERP is configured to always fallback to WebSockets"
|
||||
} else {
|
||||
for k, v := range c.lastDERPForcedWebSockets {
|
||||
derpForcedWebsocket[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
node := &Node{
|
||||
ID: c.netMap.SelfNode.ID,
|
||||
AsOf: dbtime.Now(),
|
||||
Key: c.netMap.SelfNode.Key,
|
||||
Addresses: c.netMap.SelfNode.Addresses,
|
||||
AllowedIPs: c.netMap.SelfNode.AllowedIPs,
|
||||
DiscoKey: c.magicConn.DiscoPublicKey(),
|
||||
Endpoints: endpoints,
|
||||
PreferredDERP: preferredDERP,
|
||||
DERPLatency: derpLatency,
|
||||
DERPForcedWebsocket: derpForcedWebsocket,
|
||||
}
|
||||
c.mutex.Lock()
|
||||
if c.blockEndpoints {
|
||||
node.Endpoints = nil
|
||||
}
|
||||
c.mutex.Unlock()
|
||||
return node
|
||||
c.nodeUpdater.L.Lock()
|
||||
defer c.nodeUpdater.L.Unlock()
|
||||
return c.nodeUpdater.nodeLocked()
|
||||
}
|
||||
|
||||
// This and below is taken _mostly_ verbatim from Tailscale:
|
||||
@@ -1056,15 +673,3 @@ func Logger(logger slog.Logger) tslogger.Logf {
|
||||
logger.Debug(context.Background(), fmt.Sprintf(format, args...))
|
||||
})
|
||||
}
|
||||
|
||||
func endpointsEqual(x, y []tailcfg.Endpoint) bool {
|
||||
if len(x) != len(y) {
|
||||
return false
|
||||
}
|
||||
for i := range x {
|
||||
if x[i] != y[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
+54
-29
@@ -5,6 +5,7 @@ import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/coder/v2/tailnet/tailnettest"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
@@ -22,10 +24,10 @@ func TestMain(m *testing.M) {
|
||||
|
||||
func TestTailnet(t *testing.T) {
|
||||
t.Parallel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
derpMap, _ := tailnettest.RunDERPAndSTUN(t)
|
||||
t.Run("InstantClose", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
conn, err := tailnet.NewConn(&tailnet.Options{
|
||||
Addresses: []netip.Prefix{netip.PrefixFrom(tailnet.IP(), 128)},
|
||||
Logger: logger.Named("w1"),
|
||||
@@ -37,6 +39,8 @@ func TestTailnet(t *testing.T) {
|
||||
})
|
||||
t.Run("Connect", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
w1IP := tailnet.IP()
|
||||
w1, err := tailnet.NewConn(&tailnet.Options{
|
||||
Addresses: []netip.Prefix{netip.PrefixFrom(w1IP, 128)},
|
||||
@@ -55,14 +59,8 @@ func TestTailnet(t *testing.T) {
|
||||
_ = w1.Close()
|
||||
_ = w2.Close()
|
||||
})
|
||||
w1.SetNodeCallback(func(node *tailnet.Node) {
|
||||
err := w2.UpdateNodes([]*tailnet.Node{node}, false)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
w2.SetNodeCallback(func(node *tailnet.Node) {
|
||||
err := w1.UpdateNodes([]*tailnet.Node{node}, false)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
stitch(t, w2, w1)
|
||||
stitch(t, w1, w2)
|
||||
require.True(t, w2.AwaitReachable(context.Background(), w1IP))
|
||||
conn := make(chan struct{}, 1)
|
||||
go func() {
|
||||
@@ -89,7 +87,7 @@ func TestTailnet(t *testing.T) {
|
||||
default:
|
||||
}
|
||||
})
|
||||
node := <-nodes
|
||||
node := testutil.RequireRecvCtx(ctx, t, nodes)
|
||||
// Ensure this connected over DERP!
|
||||
require.Len(t, node.DERPForcedWebsocket, 0)
|
||||
|
||||
@@ -99,6 +97,7 @@ func TestTailnet(t *testing.T) {
|
||||
|
||||
t.Run("ForcesWebSockets", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
w1IP := tailnet.IP()
|
||||
@@ -122,14 +121,8 @@ func TestTailnet(t *testing.T) {
|
||||
_ = w1.Close()
|
||||
_ = w2.Close()
|
||||
})
|
||||
w1.SetNodeCallback(func(node *tailnet.Node) {
|
||||
err := w2.UpdateNodes([]*tailnet.Node{node}, false)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
w2.SetNodeCallback(func(node *tailnet.Node) {
|
||||
err := w1.UpdateNodes([]*tailnet.Node{node}, false)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
stitch(t, w2, w1)
|
||||
stitch(t, w1, w2)
|
||||
require.True(t, w2.AwaitReachable(ctx, w1IP))
|
||||
conn := make(chan struct{}, 1)
|
||||
go func() {
|
||||
@@ -243,11 +236,16 @@ func TestConn_UpdateDERP(t *testing.T) {
|
||||
err := client1.Close()
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
client1.SetNodeCallback(func(node *tailnet.Node) {
|
||||
err := conn.UpdateNodes([]*tailnet.Node{node}, false)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
client1.UpdateNodes([]*tailnet.Node{conn.Node()}, false)
|
||||
stitch(t, conn, client1)
|
||||
pn, err := tailnet.NodeToProto(conn.Node())
|
||||
require.NoError(t, err)
|
||||
connID := uuid.New()
|
||||
err = client1.UpdatePeers([]*proto.CoordinateResponse_PeerUpdate{{
|
||||
Id: connID[:],
|
||||
Node: pn,
|
||||
Kind: proto.CoordinateResponse_PeerUpdate_NODE,
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
awaitReachableCtx1, awaitReachableCancel1 := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer awaitReachableCancel1()
|
||||
@@ -288,7 +286,13 @@ parentLoop:
|
||||
|
||||
// ... unless the client updates it's derp map and nodes.
|
||||
client1.SetDERPMap(derpMap2)
|
||||
client1.UpdateNodes([]*tailnet.Node{conn.Node()}, false)
|
||||
pn, err = tailnet.NodeToProto(conn.Node())
|
||||
require.NoError(t, err)
|
||||
client1.UpdatePeers([]*proto.CoordinateResponse_PeerUpdate{{
|
||||
Id: connID[:],
|
||||
Node: pn,
|
||||
Kind: proto.CoordinateResponse_PeerUpdate_NODE,
|
||||
}})
|
||||
awaitReachableCtx3, awaitReachableCancel3 := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer awaitReachableCancel3()
|
||||
require.True(t, client1.AwaitReachable(awaitReachableCtx3, ip))
|
||||
@@ -306,13 +310,34 @@ parentLoop:
|
||||
err := client2.Close()
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
client2.SetNodeCallback(func(node *tailnet.Node) {
|
||||
err := conn.UpdateNodes([]*tailnet.Node{node}, false)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
client2.UpdateNodes([]*tailnet.Node{conn.Node()}, false)
|
||||
stitch(t, conn, client2)
|
||||
pn, err = tailnet.NodeToProto(conn.Node())
|
||||
require.NoError(t, err)
|
||||
client2.UpdatePeers([]*proto.CoordinateResponse_PeerUpdate{{
|
||||
Id: connID[:],
|
||||
Node: pn,
|
||||
Kind: proto.CoordinateResponse_PeerUpdate_NODE,
|
||||
}})
|
||||
|
||||
awaitReachableCtx4, awaitReachableCancel4 := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer awaitReachableCancel4()
|
||||
require.True(t, client2.AwaitReachable(awaitReachableCtx4, ip))
|
||||
}
|
||||
|
||||
// stitch sends node updates from src Conn as peer updates to dst Conn. Sort of
|
||||
// like the Coordinator would, but without actually needing a Coordinator.
|
||||
func stitch(t *testing.T, dst, src *tailnet.Conn) {
|
||||
srcID := uuid.New()
|
||||
src.SetNodeCallback(func(node *tailnet.Node) {
|
||||
pn, err := tailnet.NodeToProto(node)
|
||||
if !assert.NoError(t, err) {
|
||||
return
|
||||
}
|
||||
err = dst.UpdatePeers([]*proto.CoordinateResponse_PeerUpdate{{
|
||||
Id: srcID[:],
|
||||
Node: pn,
|
||||
Kind: proto.CoordinateResponse_PeerUpdate_NODE,
|
||||
}})
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
+238
-21
@@ -3,6 +3,7 @@ package tailnet
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"html/template"
|
||||
"io"
|
||||
"net"
|
||||
@@ -92,6 +93,237 @@ type Node struct {
|
||||
Endpoints []string `json:"endpoints"`
|
||||
}
|
||||
|
||||
// Coordinatee is something that can be coordinated over the Coordinate protocol. Usually this is a
|
||||
// Conn.
|
||||
type Coordinatee interface {
|
||||
UpdatePeers([]*proto.CoordinateResponse_PeerUpdate) error
|
||||
SetNodeCallback(func(*Node))
|
||||
}
|
||||
|
||||
type Coordination interface {
|
||||
io.Closer
|
||||
Error() <-chan error
|
||||
}
|
||||
|
||||
type remoteCoordination struct {
|
||||
sync.Mutex
|
||||
closed bool
|
||||
errChan chan error
|
||||
coordinatee Coordinatee
|
||||
logger slog.Logger
|
||||
protocol proto.DRPCTailnet_CoordinateClient
|
||||
}
|
||||
|
||||
func (c *remoteCoordination) Close() error {
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
if c.closed {
|
||||
return nil
|
||||
}
|
||||
c.closed = true
|
||||
err := c.protocol.Send(&proto.CoordinateRequest{Disconnect: &proto.CoordinateRequest_Disconnect{}})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("send disconnect: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *remoteCoordination) Error() <-chan error {
|
||||
return c.errChan
|
||||
}
|
||||
|
||||
func (c *remoteCoordination) sendErr(err error) {
|
||||
select {
|
||||
case c.errChan <- err:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (c *remoteCoordination) respLoop() {
|
||||
for {
|
||||
resp, err := c.protocol.Recv()
|
||||
if err != nil {
|
||||
c.sendErr(xerrors.Errorf("read: %w", err))
|
||||
return
|
||||
}
|
||||
err = c.coordinatee.UpdatePeers(resp.GetPeerUpdates())
|
||||
if err != nil {
|
||||
c.sendErr(xerrors.Errorf("update peers: %w", err))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// NewRemoteCoordination uses the provided protocol to coordinate the provided coordinee (usually a
|
||||
// Conn). If the tunnelTarget is not uuid.Nil, then we add a tunnel to the peer (i.e. we are acting as
|
||||
// a client---agents should NOT set this!).
|
||||
func NewRemoteCoordination(logger slog.Logger,
|
||||
protocol proto.DRPCTailnet_CoordinateClient, coordinatee Coordinatee,
|
||||
tunnelTarget uuid.UUID,
|
||||
) Coordination {
|
||||
c := &remoteCoordination{
|
||||
errChan: make(chan error, 1),
|
||||
coordinatee: coordinatee,
|
||||
logger: logger,
|
||||
protocol: protocol,
|
||||
}
|
||||
if tunnelTarget != uuid.Nil {
|
||||
c.Lock()
|
||||
err := c.protocol.Send(&proto.CoordinateRequest{AddTunnel: &proto.CoordinateRequest_Tunnel{Id: tunnelTarget[:]}})
|
||||
c.Unlock()
|
||||
if err != nil {
|
||||
c.sendErr(err)
|
||||
}
|
||||
}
|
||||
|
||||
coordinatee.SetNodeCallback(func(node *Node) {
|
||||
pn, err := NodeToProto(node)
|
||||
if err != nil {
|
||||
c.logger.Critical(context.Background(), "failed to convert node", slog.Error(err))
|
||||
c.sendErr(err)
|
||||
return
|
||||
}
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
if c.closed {
|
||||
c.logger.Debug(context.Background(), "ignored node update because coordination is closed")
|
||||
return
|
||||
}
|
||||
err = c.protocol.Send(&proto.CoordinateRequest{UpdateSelf: &proto.CoordinateRequest_UpdateSelf{Node: pn}})
|
||||
if err != nil {
|
||||
c.sendErr(xerrors.Errorf("write: %w", err))
|
||||
}
|
||||
})
|
||||
go c.respLoop()
|
||||
return c
|
||||
}
|
||||
|
||||
type inMemoryCoordination struct {
|
||||
sync.Mutex
|
||||
ctx context.Context
|
||||
errChan chan error
|
||||
closed bool
|
||||
closedCh chan struct{}
|
||||
coordinatee Coordinatee
|
||||
logger slog.Logger
|
||||
resps <-chan *proto.CoordinateResponse
|
||||
reqs chan<- *proto.CoordinateRequest
|
||||
}
|
||||
|
||||
func (c *inMemoryCoordination) sendErr(err error) {
|
||||
select {
|
||||
case c.errChan <- err:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (c *inMemoryCoordination) Error() <-chan error {
|
||||
return c.errChan
|
||||
}
|
||||
|
||||
// NewInMemoryCoordination connects a Coordinatee (usually Conn) to an in memory Coordinator, for testing
|
||||
// or local clients. Set ClientID to uuid.Nil for an agent.
|
||||
func NewInMemoryCoordination(
|
||||
ctx context.Context, logger slog.Logger,
|
||||
clientID, agentID uuid.UUID,
|
||||
coordinator Coordinator, coordinatee Coordinatee,
|
||||
) Coordination {
|
||||
thisID := agentID
|
||||
logger = logger.With(slog.F("agent_id", agentID))
|
||||
var auth TunnelAuth = AgentTunnelAuth{}
|
||||
if clientID != uuid.Nil {
|
||||
// this is a client connection
|
||||
auth = ClientTunnelAuth{AgentID: agentID}
|
||||
logger = logger.With(slog.F("client_id", clientID))
|
||||
thisID = clientID
|
||||
}
|
||||
c := &inMemoryCoordination{
|
||||
ctx: ctx,
|
||||
errChan: make(chan error, 1),
|
||||
coordinatee: coordinatee,
|
||||
logger: logger,
|
||||
closedCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
// use the background context since we will depend exclusively on closing the req channel to
|
||||
// tell the coordinator we are done.
|
||||
c.reqs, c.resps = coordinator.Coordinate(context.Background(),
|
||||
thisID, fmt.Sprintf("inmemory%s", thisID),
|
||||
auth,
|
||||
)
|
||||
go c.respLoop()
|
||||
if agentID != uuid.Nil {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
c.logger.Warn(ctx, "context expired before we could add tunnel", slog.Error(ctx.Err()))
|
||||
return c
|
||||
case c.reqs <- &proto.CoordinateRequest{AddTunnel: &proto.CoordinateRequest_Tunnel{Id: agentID[:]}}:
|
||||
// OK!
|
||||
}
|
||||
}
|
||||
coordinatee.SetNodeCallback(func(n *Node) {
|
||||
pn, err := NodeToProto(n)
|
||||
if err != nil {
|
||||
c.logger.Critical(ctx, "failed to convert node", slog.Error(err))
|
||||
c.sendErr(err)
|
||||
return
|
||||
}
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
if c.closed {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
c.logger.Info(ctx, "context expired before sending node update")
|
||||
return
|
||||
case c.reqs <- &proto.CoordinateRequest{UpdateSelf: &proto.CoordinateRequest_UpdateSelf{Node: pn}}:
|
||||
c.logger.Debug(ctx, "sent node in-memory to coordinator")
|
||||
}
|
||||
})
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *inMemoryCoordination) respLoop() {
|
||||
for {
|
||||
select {
|
||||
case <-c.closedCh:
|
||||
c.logger.Debug(context.Background(), "in-memory coordination closed")
|
||||
return
|
||||
case resp, ok := <-c.resps:
|
||||
if !ok {
|
||||
c.logger.Debug(context.Background(), "in-memory response channel closed")
|
||||
return
|
||||
}
|
||||
c.logger.Debug(context.Background(), "got in-memory response from coordinator", slog.F("resp", resp))
|
||||
err := c.coordinatee.UpdatePeers(resp.GetPeerUpdates())
|
||||
if err != nil {
|
||||
c.sendErr(xerrors.Errorf("failed to update peers: %w", err))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *inMemoryCoordination) Close() error {
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
c.logger.Debug(context.Background(), "closing in-memory coordination")
|
||||
if c.closed {
|
||||
return nil
|
||||
}
|
||||
defer close(c.reqs)
|
||||
c.closed = true
|
||||
close(c.closedCh)
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return xerrors.Errorf("failed to gracefully disconnect: %w", c.ctx.Err())
|
||||
case c.reqs <- &proto.CoordinateRequest{Disconnect: &proto.CoordinateRequest_Disconnect{}}:
|
||||
c.logger.Debug(context.Background(), "sent graceful disconnect in-memory")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// ServeCoordinator matches the RW structure of a coordinator to exchange node messages.
|
||||
func ServeCoordinator(conn net.Conn, updateNodes func(node []*Node) error) (func(node *Node), <-chan error) {
|
||||
errChan := make(chan error, 1)
|
||||
@@ -237,21 +469,17 @@ func ServeMultiAgent(c CoordinatorV2, logger slog.Logger, id uuid.UUID) MultiAge
|
||||
}
|
||||
return false
|
||||
},
|
||||
OnSubscribe: func(enq Queue, agent uuid.UUID) (*Node, error) {
|
||||
OnSubscribe: func(enq Queue, agent uuid.UUID) error {
|
||||
err := SendCtx(ctx, reqs, &proto.CoordinateRequest{AddTunnel: &proto.CoordinateRequest_Tunnel{Id: UUIDToByteSlice(agent)}})
|
||||
return c.Node(agent), err
|
||||
return err
|
||||
},
|
||||
OnUnsubscribe: func(enq Queue, agent uuid.UUID) error {
|
||||
err := SendCtx(ctx, reqs, &proto.CoordinateRequest{RemoveTunnel: &proto.CoordinateRequest_Tunnel{Id: UUIDToByteSlice(agent)}})
|
||||
return err
|
||||
},
|
||||
OnNodeUpdate: func(id uuid.UUID, node *Node) error {
|
||||
pn, err := NodeToProto(node)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
OnNodeUpdate: func(id uuid.UUID, node *proto.Node) error {
|
||||
return SendCtx(ctx, reqs, &proto.CoordinateRequest{UpdateSelf: &proto.CoordinateRequest_UpdateSelf{
|
||||
Node: pn,
|
||||
Node: node,
|
||||
}})
|
||||
},
|
||||
OnRemove: func(_ Queue) {
|
||||
@@ -285,7 +513,7 @@ const (
|
||||
type Queue interface {
|
||||
UniqueID() uuid.UUID
|
||||
Kind() QueueKind
|
||||
Enqueue(n []*Node) error
|
||||
Enqueue(resp *proto.CoordinateResponse) error
|
||||
Name() string
|
||||
Stats() (start, lastWrite int64)
|
||||
Overwrites() int64
|
||||
@@ -793,18 +1021,7 @@ func v1RespLoop(ctx context.Context, cancel context.CancelFunc, logger slog.Logg
|
||||
return
|
||||
}
|
||||
logger.Debug(ctx, "v1RespLoop got response", slog.F("resp", resp))
|
||||
nodes, err := OnlyNodeUpdates(resp)
|
||||
if err != nil {
|
||||
logger.Critical(ctx, "v1RespLoop failed to decode resp", slog.F("resp", resp), slog.Error(err))
|
||||
_ = q.CoordinatorClose()
|
||||
return
|
||||
}
|
||||
// don't send empty updates
|
||||
if len(nodes) == 0 {
|
||||
logger.Debug(ctx, "v1RespLoop skipping enqueueing 0-length v1 update")
|
||||
continue
|
||||
}
|
||||
err = q.Enqueue(nodes)
|
||||
err = q.Enqueue(resp)
|
||||
if err != nil && !xerrors.Is(err, context.Canceled) {
|
||||
logger.Error(ctx, "v1RespLoop failed to enqueue v1 update", slog.Error(err))
|
||||
}
|
||||
|
||||
+17
-19
@@ -8,13 +8,15 @@ import (
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
)
|
||||
|
||||
type MultiAgentConn interface {
|
||||
UpdateSelf(node *Node) error
|
||||
UpdateSelf(node *proto.Node) error
|
||||
SubscribeAgent(agentID uuid.UUID) error
|
||||
UnsubscribeAgent(agentID uuid.UUID) error
|
||||
NextUpdate(ctx context.Context) ([]*Node, bool)
|
||||
NextUpdate(ctx context.Context) (*proto.CoordinateResponse, bool)
|
||||
AgentIsLegacy(agentID uuid.UUID) bool
|
||||
Close() error
|
||||
IsClosed() bool
|
||||
@@ -26,16 +28,16 @@ type MultiAgent struct {
|
||||
ID uuid.UUID
|
||||
|
||||
AgentIsLegacyFunc func(agentID uuid.UUID) bool
|
||||
OnSubscribe func(enq Queue, agent uuid.UUID) (*Node, error)
|
||||
OnSubscribe func(enq Queue, agent uuid.UUID) error
|
||||
OnUnsubscribe func(enq Queue, agent uuid.UUID) error
|
||||
OnNodeUpdate func(id uuid.UUID, node *Node) error
|
||||
OnNodeUpdate func(id uuid.UUID, node *proto.Node) error
|
||||
OnRemove func(enq Queue)
|
||||
|
||||
ctx context.Context
|
||||
ctxCancel func()
|
||||
closed bool
|
||||
|
||||
updates chan []*Node
|
||||
updates chan *proto.CoordinateResponse
|
||||
closeOnce sync.Once
|
||||
start int64
|
||||
lastWrite int64
|
||||
@@ -45,7 +47,7 @@ type MultiAgent struct {
|
||||
}
|
||||
|
||||
func (m *MultiAgent) Init() *MultiAgent {
|
||||
m.updates = make(chan []*Node, 128)
|
||||
m.updates = make(chan *proto.CoordinateResponse, 128)
|
||||
m.start = time.Now().Unix()
|
||||
m.ctx, m.ctxCancel = context.WithCancel(context.Background())
|
||||
return m
|
||||
@@ -65,7 +67,7 @@ func (m *MultiAgent) AgentIsLegacy(agentID uuid.UUID) bool {
|
||||
|
||||
var ErrMultiAgentClosed = xerrors.New("multiagent is closed")
|
||||
|
||||
func (m *MultiAgent) UpdateSelf(node *Node) error {
|
||||
func (m *MultiAgent) UpdateSelf(node *proto.Node) error {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
if m.closed {
|
||||
@@ -82,15 +84,11 @@ func (m *MultiAgent) SubscribeAgent(agentID uuid.UUID) error {
|
||||
return ErrMultiAgentClosed
|
||||
}
|
||||
|
||||
node, err := m.OnSubscribe(m, agentID)
|
||||
err := m.OnSubscribe(m, agentID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if node != nil {
|
||||
return m.enqueueLocked([]*Node{node})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -104,17 +102,17 @@ func (m *MultiAgent) UnsubscribeAgent(agentID uuid.UUID) error {
|
||||
return m.OnUnsubscribe(m, agentID)
|
||||
}
|
||||
|
||||
func (m *MultiAgent) NextUpdate(ctx context.Context) ([]*Node, bool) {
|
||||
func (m *MultiAgent) NextUpdate(ctx context.Context) (*proto.CoordinateResponse, bool) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, false
|
||||
|
||||
case nodes, ok := <-m.updates:
|
||||
return nodes, ok
|
||||
case resp, ok := <-m.updates:
|
||||
return resp, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (m *MultiAgent) Enqueue(nodes []*Node) error {
|
||||
func (m *MultiAgent) Enqueue(resp *proto.CoordinateResponse) error {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
@@ -122,14 +120,14 @@ func (m *MultiAgent) Enqueue(nodes []*Node) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
return m.enqueueLocked(nodes)
|
||||
return m.enqueueLocked(resp)
|
||||
}
|
||||
|
||||
func (m *MultiAgent) enqueueLocked(nodes []*Node) error {
|
||||
func (m *MultiAgent) enqueueLocked(resp *proto.CoordinateResponse) error {
|
||||
atomic.StoreInt64(&m.lastWrite, time.Now().Unix())
|
||||
|
||||
select {
|
||||
case m.updates <- nodes:
|
||||
case m.updates <- resp:
|
||||
return nil
|
||||
default:
|
||||
return ErrWouldBlock
|
||||
|
||||
+3
-1
@@ -75,7 +75,9 @@ func NewClientService(
|
||||
}
|
||||
server := drpcserver.NewWithOptions(mux, drpcserver.Options{
|
||||
Log: func(err error) {
|
||||
if xerrors.Is(err, io.EOF) {
|
||||
if xerrors.Is(err, io.EOF) ||
|
||||
xerrors.Is(err, context.Canceled) ||
|
||||
xerrors.Is(err, context.DeadlineExceeded) {
|
||||
return
|
||||
}
|
||||
logger.Debug(context.Background(), "drpc server error", slog.Error(err))
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: github.com/coder/coder/v2/tailnet (interfaces: Coordinator)
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -destination ./coordinatormock.go -package tailnettest github.com/coder/coder/v2/tailnet Coordinator
|
||||
//
|
||||
|
||||
// Package tailnettest is a generated GoMock package.
|
||||
package tailnettest
|
||||
|
||||
import (
|
||||
context "context"
|
||||
net "net"
|
||||
http "net/http"
|
||||
reflect "reflect"
|
||||
|
||||
tailnet "github.com/coder/coder/v2/tailnet"
|
||||
proto "github.com/coder/coder/v2/tailnet/proto"
|
||||
uuid "github.com/google/uuid"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockCoordinator is a mock of Coordinator interface.
|
||||
type MockCoordinator struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockCoordinatorMockRecorder
|
||||
}
|
||||
|
||||
// MockCoordinatorMockRecorder is the mock recorder for MockCoordinator.
|
||||
type MockCoordinatorMockRecorder struct {
|
||||
mock *MockCoordinator
|
||||
}
|
||||
|
||||
// NewMockCoordinator creates a new mock instance.
|
||||
func NewMockCoordinator(ctrl *gomock.Controller) *MockCoordinator {
|
||||
mock := &MockCoordinator{ctrl: ctrl}
|
||||
mock.recorder = &MockCoordinatorMockRecorder{mock}
|
||||
return mock
|
||||
}
|
||||
|
||||
// EXPECT returns an object that allows the caller to indicate expected use.
|
||||
func (m *MockCoordinator) EXPECT() *MockCoordinatorMockRecorder {
|
||||
return m.recorder
|
||||
}
|
||||
|
||||
// Close mocks base method.
|
||||
func (m *MockCoordinator) Close() error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Close")
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// Close indicates an expected call of Close.
|
||||
func (mr *MockCoordinatorMockRecorder) Close() *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockCoordinator)(nil).Close))
|
||||
}
|
||||
|
||||
// Coordinate mocks base method.
|
||||
func (m *MockCoordinator) Coordinate(arg0 context.Context, arg1 uuid.UUID, arg2 string, arg3 tailnet.TunnelAuth) (chan<- *proto.CoordinateRequest, <-chan *proto.CoordinateResponse) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Coordinate", arg0, arg1, arg2, arg3)
|
||||
ret0, _ := ret[0].(chan<- *proto.CoordinateRequest)
|
||||
ret1, _ := ret[1].(<-chan *proto.CoordinateResponse)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// Coordinate indicates an expected call of Coordinate.
|
||||
func (mr *MockCoordinatorMockRecorder) Coordinate(arg0, arg1, arg2, arg3 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Coordinate", reflect.TypeOf((*MockCoordinator)(nil).Coordinate), arg0, arg1, arg2, arg3)
|
||||
}
|
||||
|
||||
// Node mocks base method.
|
||||
func (m *MockCoordinator) Node(arg0 uuid.UUID) *tailnet.Node {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Node", arg0)
|
||||
ret0, _ := ret[0].(*tailnet.Node)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// Node indicates an expected call of Node.
|
||||
func (mr *MockCoordinatorMockRecorder) Node(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Node", reflect.TypeOf((*MockCoordinator)(nil).Node), arg0)
|
||||
}
|
||||
|
||||
// ServeAgent mocks base method.
|
||||
func (m *MockCoordinator) ServeAgent(arg0 net.Conn, arg1 uuid.UUID, arg2 string) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ServeAgent", arg0, arg1, arg2)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ServeAgent indicates an expected call of ServeAgent.
|
||||
func (mr *MockCoordinatorMockRecorder) ServeAgent(arg0, arg1, arg2 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ServeAgent", reflect.TypeOf((*MockCoordinator)(nil).ServeAgent), arg0, arg1, arg2)
|
||||
}
|
||||
|
||||
// ServeClient mocks base method.
|
||||
func (m *MockCoordinator) ServeClient(arg0 net.Conn, arg1, arg2 uuid.UUID) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ServeClient", arg0, arg1, arg2)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ServeClient indicates an expected call of ServeClient.
|
||||
func (mr *MockCoordinatorMockRecorder) ServeClient(arg0, arg1, arg2 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ServeClient", reflect.TypeOf((*MockCoordinator)(nil).ServeClient), arg0, arg1, arg2)
|
||||
}
|
||||
|
||||
// ServeHTTPDebug mocks base method.
|
||||
func (m *MockCoordinator) ServeHTTPDebug(arg0 http.ResponseWriter, arg1 *http.Request) {
|
||||
m.ctrl.T.Helper()
|
||||
m.ctrl.Call(m, "ServeHTTPDebug", arg0, arg1)
|
||||
}
|
||||
|
||||
// ServeHTTPDebug indicates an expected call of ServeHTTPDebug.
|
||||
func (mr *MockCoordinatorMockRecorder) ServeHTTPDebug(arg0, arg1 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ServeHTTPDebug", reflect.TypeOf((*MockCoordinator)(nil).ServeHTTPDebug), arg0, arg1)
|
||||
}
|
||||
|
||||
// ServeMultiAgent mocks base method.
|
||||
func (m *MockCoordinator) ServeMultiAgent(arg0 uuid.UUID) tailnet.MultiAgentConn {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ServeMultiAgent", arg0)
|
||||
ret0, _ := ret[0].(tailnet.MultiAgentConn)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ServeMultiAgent indicates an expected call of ServeMultiAgent.
|
||||
func (mr *MockCoordinatorMockRecorder) ServeMultiAgent(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ServeMultiAgent", reflect.TypeOf((*MockCoordinator)(nil).ServeMultiAgent), arg0)
|
||||
}
|
||||
@@ -1,141 +0,0 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: github.com/coder/coder/v2/tailnet (interfaces: MultiAgentConn)
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -destination ./multiagentmock.go -package tailnettest github.com/coder/coder/v2/tailnet MultiAgentConn
|
||||
//
|
||||
|
||||
// Package tailnettest is a generated GoMock package.
|
||||
package tailnettest
|
||||
|
||||
import (
|
||||
context "context"
|
||||
reflect "reflect"
|
||||
|
||||
tailnet "github.com/coder/coder/v2/tailnet"
|
||||
uuid "github.com/google/uuid"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockMultiAgentConn is a mock of MultiAgentConn interface.
|
||||
type MockMultiAgentConn struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockMultiAgentConnMockRecorder
|
||||
}
|
||||
|
||||
// MockMultiAgentConnMockRecorder is the mock recorder for MockMultiAgentConn.
|
||||
type MockMultiAgentConnMockRecorder struct {
|
||||
mock *MockMultiAgentConn
|
||||
}
|
||||
|
||||
// NewMockMultiAgentConn creates a new mock instance.
|
||||
func NewMockMultiAgentConn(ctrl *gomock.Controller) *MockMultiAgentConn {
|
||||
mock := &MockMultiAgentConn{ctrl: ctrl}
|
||||
mock.recorder = &MockMultiAgentConnMockRecorder{mock}
|
||||
return mock
|
||||
}
|
||||
|
||||
// EXPECT returns an object that allows the caller to indicate expected use.
|
||||
func (m *MockMultiAgentConn) EXPECT() *MockMultiAgentConnMockRecorder {
|
||||
return m.recorder
|
||||
}
|
||||
|
||||
// AgentIsLegacy mocks base method.
|
||||
func (m *MockMultiAgentConn) AgentIsLegacy(arg0 uuid.UUID) bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "AgentIsLegacy", arg0)
|
||||
ret0, _ := ret[0].(bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// AgentIsLegacy indicates an expected call of AgentIsLegacy.
|
||||
func (mr *MockMultiAgentConnMockRecorder) AgentIsLegacy(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AgentIsLegacy", reflect.TypeOf((*MockMultiAgentConn)(nil).AgentIsLegacy), arg0)
|
||||
}
|
||||
|
||||
// Close mocks base method.
|
||||
func (m *MockMultiAgentConn) Close() error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Close")
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// Close indicates an expected call of Close.
|
||||
func (mr *MockMultiAgentConnMockRecorder) Close() *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Close", reflect.TypeOf((*MockMultiAgentConn)(nil).Close))
|
||||
}
|
||||
|
||||
// IsClosed mocks base method.
|
||||
func (m *MockMultiAgentConn) IsClosed() bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "IsClosed")
|
||||
ret0, _ := ret[0].(bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// IsClosed indicates an expected call of IsClosed.
|
||||
func (mr *MockMultiAgentConnMockRecorder) IsClosed() *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "IsClosed", reflect.TypeOf((*MockMultiAgentConn)(nil).IsClosed))
|
||||
}
|
||||
|
||||
// NextUpdate mocks base method.
|
||||
func (m *MockMultiAgentConn) NextUpdate(arg0 context.Context) ([]*tailnet.Node, bool) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "NextUpdate", arg0)
|
||||
ret0, _ := ret[0].([]*tailnet.Node)
|
||||
ret1, _ := ret[1].(bool)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// NextUpdate indicates an expected call of NextUpdate.
|
||||
func (mr *MockMultiAgentConnMockRecorder) NextUpdate(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "NextUpdate", reflect.TypeOf((*MockMultiAgentConn)(nil).NextUpdate), arg0)
|
||||
}
|
||||
|
||||
// SubscribeAgent mocks base method.
|
||||
func (m *MockMultiAgentConn) SubscribeAgent(arg0 uuid.UUID) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "SubscribeAgent", arg0)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// SubscribeAgent indicates an expected call of SubscribeAgent.
|
||||
func (mr *MockMultiAgentConnMockRecorder) SubscribeAgent(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SubscribeAgent", reflect.TypeOf((*MockMultiAgentConn)(nil).SubscribeAgent), arg0)
|
||||
}
|
||||
|
||||
// UnsubscribeAgent mocks base method.
|
||||
func (m *MockMultiAgentConn) UnsubscribeAgent(arg0 uuid.UUID) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "UnsubscribeAgent", arg0)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// UnsubscribeAgent indicates an expected call of UnsubscribeAgent.
|
||||
func (mr *MockMultiAgentConnMockRecorder) UnsubscribeAgent(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnsubscribeAgent", reflect.TypeOf((*MockMultiAgentConn)(nil).UnsubscribeAgent), arg0)
|
||||
}
|
||||
|
||||
// UpdateSelf mocks base method.
|
||||
func (m *MockMultiAgentConn) UpdateSelf(arg0 *tailnet.Node) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "UpdateSelf", arg0)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// UpdateSelf indicates an expected call of UpdateSelf.
|
||||
func (mr *MockMultiAgentConnMockRecorder) UpdateSelf(arg0 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateSelf", reflect.TypeOf((*MockMultiAgentConn)(nil).UpdateSelf), arg0)
|
||||
}
|
||||
@@ -21,7 +21,7 @@ import (
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
)
|
||||
|
||||
//go:generate mockgen -destination ./multiagentmock.go -package tailnettest github.com/coder/coder/v2/tailnet MultiAgentConn
|
||||
//go:generate mockgen -destination ./coordinatormock.go -package tailnettest github.com/coder/coder/v2/tailnet Coordinator
|
||||
|
||||
// RunDERPAndSTUN creates a DERP mapping for tests.
|
||||
func RunDERPAndSTUN(t *testing.T) (*tailcfg.DERPMap, *derp.Server) {
|
||||
|
||||
+15
-5
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -29,7 +30,7 @@ type TrackedConn struct {
|
||||
cancel func()
|
||||
kind QueueKind
|
||||
conn net.Conn
|
||||
updates chan []*Node
|
||||
updates chan *proto.CoordinateResponse
|
||||
logger slog.Logger
|
||||
lastData []byte
|
||||
|
||||
@@ -55,7 +56,7 @@ func NewTrackedConn(ctx context.Context, cancel func(),
|
||||
// coordinator mutex while queuing. Node updates don't
|
||||
// come quickly, so 512 should be plenty for all but
|
||||
// the most pathological cases.
|
||||
updates := make(chan []*Node, ResponseBufferSize)
|
||||
updates := make(chan *proto.CoordinateResponse, ResponseBufferSize)
|
||||
now := time.Now().Unix()
|
||||
return &TrackedConn{
|
||||
ctx: ctx,
|
||||
@@ -72,10 +73,10 @@ func NewTrackedConn(ctx context.Context, cancel func(),
|
||||
}
|
||||
}
|
||||
|
||||
func (t *TrackedConn) Enqueue(n []*Node) (err error) {
|
||||
func (t *TrackedConn) Enqueue(resp *proto.CoordinateResponse) (err error) {
|
||||
atomic.StoreInt64(&t.lastWrite, time.Now().Unix())
|
||||
select {
|
||||
case t.updates <- n:
|
||||
case t.updates <- resp:
|
||||
return nil
|
||||
default:
|
||||
return ErrWouldBlock
|
||||
@@ -124,7 +125,16 @@ func (t *TrackedConn) SendUpdates() {
|
||||
case <-t.ctx.Done():
|
||||
t.logger.Debug(t.ctx, "done sending updates")
|
||||
return
|
||||
case nodes := <-t.updates:
|
||||
case resp := <-t.updates:
|
||||
nodes, err := OnlyNodeUpdates(resp)
|
||||
if err != nil {
|
||||
t.logger.Critical(t.ctx, "unable to parse response", slog.Error(err))
|
||||
return
|
||||
}
|
||||
if len(nodes) == 0 {
|
||||
t.logger.Debug(t.ctx, "skipping response with no nodes")
|
||||
continue
|
||||
}
|
||||
data, err := json.Marshal(nodes)
|
||||
if err != nil {
|
||||
t.logger.Error(t.ctx, "unable to marshal nodes update", slog.Error(err), slog.F("nodes", nodes))
|
||||
|
||||
Reference in New Issue
Block a user