mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: implement CoderVPN client & tunnel (#15612)
Addresses #14734. This PR wires up `tunnel.go` to a `tailnet.Conn` via the new `/tailnet` endpoint, with all the necessary controllers such that a VPN connection can be started, stopped and inspected via the CoderVPN protocol.
This commit is contained in:
+17
-7
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/cenkalti/backoff/v4"
|
||||
"github.com/google/uuid"
|
||||
"github.com/tailscale/wireguard-go/tun"
|
||||
"golang.org/x/xerrors"
|
||||
"google.golang.org/protobuf/types/known/durationpb"
|
||||
"google.golang.org/protobuf/types/known/wrapperspb"
|
||||
@@ -113,6 +114,8 @@ type Options struct {
|
||||
DNSConfigurator dns.OSConfigurator
|
||||
// Router is optional, and is passed to the underlying wireguard engine.
|
||||
Router router.Router
|
||||
// TUNDev is optional, and is passed to the underlying wireguard engine.
|
||||
TUNDev tun.Device
|
||||
}
|
||||
|
||||
// TelemetrySink allows tailnet.Conn to send network telemetry to the Coder
|
||||
@@ -143,6 +146,8 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
return nil, xerrors.New("At least one IP range must be provided")
|
||||
}
|
||||
|
||||
netns.SetEnabled(options.TUNDev != nil)
|
||||
|
||||
var telemetryStore *TelemetryStore
|
||||
if options.TelemetrySink != nil {
|
||||
var err error
|
||||
@@ -187,6 +192,7 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
SetSubsystem: sys.Set,
|
||||
DNS: options.DNSConfigurator,
|
||||
Router: options.Router,
|
||||
Tun: options.TUNDev,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("create wgengine: %w", err)
|
||||
@@ -197,11 +203,14 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
}
|
||||
}()
|
||||
wireguardEngine.InstallCaptureHook(options.CaptureHook)
|
||||
dialer.UseNetstackForIP = func(ip netip.Addr) bool {
|
||||
_, ok := wireguardEngine.PeerForIP(ip)
|
||||
return ok
|
||||
if options.TUNDev == nil {
|
||||
dialer.UseNetstackForIP = func(ip netip.Addr) bool {
|
||||
_, ok := wireguardEngine.PeerForIP(ip)
|
||||
return ok
|
||||
}
|
||||
}
|
||||
|
||||
wireguardEngine = wgengine.NewWatchdog(wireguardEngine)
|
||||
sys.Set(wireguardEngine)
|
||||
|
||||
magicConn := sys.MagicSock.Get()
|
||||
@@ -244,11 +253,12 @@ func NewConn(options *Options) (conn *Conn, err error) {
|
||||
return nil, xerrors.Errorf("create netstack: %w", err)
|
||||
}
|
||||
|
||||
dialer.NetstackDialTCP = func(ctx context.Context, dst netip.AddrPort) (net.Conn, error) {
|
||||
return netStack.DialContextTCP(ctx, dst)
|
||||
if options.TUNDev == nil {
|
||||
dialer.NetstackDialTCP = func(ctx context.Context, dst netip.AddrPort) (net.Conn, error) {
|
||||
return netStack.DialContextTCP(ctx, dst)
|
||||
}
|
||||
netStack.ProcessLocalIPs = true
|
||||
}
|
||||
netStack.ProcessLocalIPs = true
|
||||
wireguardEngine = wgengine.NewWatchdog(wireguardEngine)
|
||||
|
||||
cfgMaps := newConfigMaps(
|
||||
options.Logger,
|
||||
|
||||
+211
-61
@@ -7,6 +7,7 @@ import (
|
||||
"maps"
|
||||
"math"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -19,6 +20,7 @@ import (
|
||||
"tailscale.com/util/dnsname"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/quartz"
|
||||
@@ -112,6 +114,11 @@ type DNSHostsSetter interface {
|
||||
SetDNSHosts(hosts map[dnsname.FQDN][]netip.Addr) error
|
||||
}
|
||||
|
||||
// UpdatesHandler is anything that expects a stream of workspace update diffs.
|
||||
type UpdatesHandler interface {
|
||||
Update(WorkspaceUpdate) error
|
||||
}
|
||||
|
||||
// ControlProtocolClients represents an abstract interface to the tailnet control plane via a set
|
||||
// of protocol clients. The Closer should close all the clients (e.g. by closing the underlying
|
||||
// connection).
|
||||
@@ -855,65 +862,121 @@ func (r *basicResumeTokenRefresher) refresh() {
|
||||
r.timer.Reset(dur, "basicResumeTokenRefresher", "refresh")
|
||||
}
|
||||
|
||||
type tunnelAllWorkspaceUpdatesController struct {
|
||||
type TunnelAllWorkspaceUpdatesController struct {
|
||||
coordCtrl *TunnelSrcCoordController
|
||||
dnsHostSetter DNSHostsSetter
|
||||
updateHandler UpdatesHandler
|
||||
ownerUsername string
|
||||
logger slog.Logger
|
||||
|
||||
mu sync.Mutex
|
||||
updater *tunnelUpdater
|
||||
}
|
||||
|
||||
type workspace struct {
|
||||
id uuid.UUID
|
||||
name string
|
||||
agents map[uuid.UUID]agent
|
||||
type Workspace struct {
|
||||
ID uuid.UUID
|
||||
Name string
|
||||
Status proto.Workspace_Status
|
||||
|
||||
ownerUsername string
|
||||
agents map[uuid.UUID]*Agent
|
||||
}
|
||||
|
||||
// addAllDNSNames adds names for all of its agents to the given map of names
|
||||
func (w workspace) addAllDNSNames(names map[dnsname.FQDN][]netip.Addr, owner string) error {
|
||||
for _, a := range w.agents {
|
||||
// updateDNSNames updates the DNS names for all agents in the workspace.
|
||||
func (w *Workspace) updateDNSNames() error {
|
||||
for id, a := range w.agents {
|
||||
names := make(map[dnsname.FQDN][]netip.Addr)
|
||||
// TODO: technically, DNS labels cannot start with numbers, but the rules are often not
|
||||
// strictly enforced.
|
||||
fqdn, err := dnsname.ToFQDN(fmt.Sprintf("%s.%s.me.coder.", a.name, w.name))
|
||||
fqdn, err := dnsname.ToFQDN(fmt.Sprintf("%s.%s.me.coder.", a.Name, w.Name))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
names[fqdn] = []netip.Addr{CoderServicePrefix.AddrFromUUID(a.id)}
|
||||
fqdn, err = dnsname.ToFQDN(fmt.Sprintf("%s.%s.%s.coder.", a.name, w.name, owner))
|
||||
names[fqdn] = []netip.Addr{CoderServicePrefix.AddrFromUUID(a.ID)}
|
||||
fqdn, err = dnsname.ToFQDN(fmt.Sprintf("%s.%s.%s.coder.", a.Name, w.Name, w.ownerUsername))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
names[fqdn] = []netip.Addr{CoderServicePrefix.AddrFromUUID(a.id)}
|
||||
}
|
||||
if len(w.agents) == 1 {
|
||||
fqdn, err := dnsname.ToFQDN(fmt.Sprintf("%s.coder.", w.name))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, a := range w.agents {
|
||||
names[fqdn] = []netip.Addr{CoderServicePrefix.AddrFromUUID(a.id)}
|
||||
names[fqdn] = []netip.Addr{CoderServicePrefix.AddrFromUUID(a.ID)}
|
||||
if len(w.agents) == 1 {
|
||||
fqdn, err := dnsname.ToFQDN(fmt.Sprintf("%s.coder.", w.Name))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, a := range w.agents {
|
||||
names[fqdn] = []netip.Addr{CoderServicePrefix.AddrFromUUID(a.ID)}
|
||||
}
|
||||
}
|
||||
a.Hosts = names
|
||||
w.agents[id] = a
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type agent struct {
|
||||
id uuid.UUID
|
||||
name string
|
||||
type Agent struct {
|
||||
ID uuid.UUID
|
||||
Name string
|
||||
WorkspaceID uuid.UUID
|
||||
Hosts map[dnsname.FQDN][]netip.Addr
|
||||
}
|
||||
|
||||
func (t *tunnelAllWorkspaceUpdatesController) New(client WorkspaceUpdatesClient) CloserWaiter {
|
||||
func (a *Agent) Clone() Agent {
|
||||
hosts := make(map[dnsname.FQDN][]netip.Addr, len(a.Hosts))
|
||||
for k, v := range a.Hosts {
|
||||
hosts[k] = slices.Clone(v)
|
||||
}
|
||||
return Agent{
|
||||
ID: a.ID,
|
||||
Name: a.Name,
|
||||
WorkspaceID: a.WorkspaceID,
|
||||
Hosts: hosts,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *TunnelAllWorkspaceUpdatesController) New(client WorkspaceUpdatesClient) CloserWaiter {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
updater := &tunnelUpdater{
|
||||
client: client,
|
||||
errChan: make(chan error, 1),
|
||||
logger: t.logger,
|
||||
coordCtrl: t.coordCtrl,
|
||||
dnsHostsSetter: t.dnsHostSetter,
|
||||
updateHandler: t.updateHandler,
|
||||
ownerUsername: t.ownerUsername,
|
||||
recvLoopDone: make(chan struct{}),
|
||||
workspaces: make(map[uuid.UUID]*workspace),
|
||||
workspaces: make(map[uuid.UUID]*Workspace),
|
||||
}
|
||||
go updater.recvLoop()
|
||||
return updater
|
||||
t.updater = updater
|
||||
go t.updater.recvLoop()
|
||||
return t.updater
|
||||
}
|
||||
|
||||
func (t *TunnelAllWorkspaceUpdatesController) CurrentState() (WorkspaceUpdate, error) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
if t.updater == nil {
|
||||
return WorkspaceUpdate{}, xerrors.New("no updater")
|
||||
}
|
||||
t.updater.Lock()
|
||||
defer t.updater.Unlock()
|
||||
out := WorkspaceUpdate{
|
||||
UpsertedWorkspaces: make([]*Workspace, 0, len(t.updater.workspaces)),
|
||||
UpsertedAgents: make([]*Agent, 0, len(t.updater.workspaces)),
|
||||
DeletedWorkspaces: make([]*Workspace, 0),
|
||||
DeletedAgents: make([]*Agent, 0),
|
||||
}
|
||||
for _, w := range t.updater.workspaces {
|
||||
out.UpsertedWorkspaces = append(out.UpsertedWorkspaces, &Workspace{
|
||||
ID: w.ID,
|
||||
Name: w.Name,
|
||||
Status: w.Status,
|
||||
})
|
||||
for _, a := range w.agents {
|
||||
out.UpsertedAgents = append(out.UpsertedAgents, ptr.Ref(a.Clone()))
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
type tunnelUpdater struct {
|
||||
@@ -922,14 +985,13 @@ type tunnelUpdater struct {
|
||||
client WorkspaceUpdatesClient
|
||||
coordCtrl *TunnelSrcCoordController
|
||||
dnsHostsSetter DNSHostsSetter
|
||||
updateHandler UpdatesHandler
|
||||
ownerUsername string
|
||||
recvLoopDone chan struct{}
|
||||
|
||||
// don't need the mutex since only manipulated by the recvLoop
|
||||
workspaces map[uuid.UUID]*workspace
|
||||
|
||||
sync.Mutex
|
||||
closed bool
|
||||
workspaces map[uuid.UUID]*Workspace
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) Close(ctx context.Context) error {
|
||||
@@ -990,18 +1052,68 @@ func (t *tunnelUpdater) recvLoop() {
|
||||
}
|
||||
}
|
||||
|
||||
type WorkspaceUpdate struct {
|
||||
UpsertedWorkspaces []*Workspace
|
||||
UpsertedAgents []*Agent
|
||||
DeletedWorkspaces []*Workspace
|
||||
DeletedAgents []*Agent
|
||||
}
|
||||
|
||||
func (w *WorkspaceUpdate) Clone() WorkspaceUpdate {
|
||||
clone := WorkspaceUpdate{
|
||||
UpsertedWorkspaces: make([]*Workspace, len(w.UpsertedWorkspaces)),
|
||||
UpsertedAgents: make([]*Agent, len(w.UpsertedAgents)),
|
||||
DeletedWorkspaces: make([]*Workspace, len(w.DeletedWorkspaces)),
|
||||
DeletedAgents: make([]*Agent, len(w.DeletedAgents)),
|
||||
}
|
||||
for i, ws := range w.UpsertedWorkspaces {
|
||||
clone.UpsertedWorkspaces[i] = &Workspace{
|
||||
ID: ws.ID,
|
||||
Name: ws.Name,
|
||||
Status: ws.Status,
|
||||
}
|
||||
}
|
||||
for i, a := range w.UpsertedAgents {
|
||||
clone.UpsertedAgents[i] = ptr.Ref(a.Clone())
|
||||
}
|
||||
for i, ws := range w.DeletedWorkspaces {
|
||||
clone.DeletedWorkspaces[i] = &Workspace{
|
||||
ID: ws.ID,
|
||||
Name: ws.Name,
|
||||
Status: ws.Status,
|
||||
}
|
||||
}
|
||||
for i, a := range w.DeletedAgents {
|
||||
clone.DeletedAgents[i] = ptr.Ref(a.Clone())
|
||||
}
|
||||
return clone
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) handleUpdate(update *proto.WorkspaceUpdate) error {
|
||||
t.Lock()
|
||||
defer t.Unlock()
|
||||
|
||||
currentUpdate := WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*Workspace{},
|
||||
UpsertedAgents: []*Agent{},
|
||||
DeletedWorkspaces: []*Workspace{},
|
||||
DeletedAgents: []*Agent{},
|
||||
}
|
||||
|
||||
for _, uw := range update.UpsertedWorkspaces {
|
||||
workspaceID, err := uuid.FromBytes(uw.Id)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to parse workspace ID: %w", err)
|
||||
}
|
||||
w := workspace{
|
||||
id: workspaceID,
|
||||
name: uw.Name,
|
||||
agents: make(map[uuid.UUID]agent),
|
||||
w := &Workspace{
|
||||
ID: workspaceID,
|
||||
Name: uw.Name,
|
||||
Status: uw.Status,
|
||||
ownerUsername: t.ownerUsername,
|
||||
agents: make(map[uuid.UUID]*Agent),
|
||||
}
|
||||
t.upsertWorkspace(w)
|
||||
t.upsertWorkspaceLocked(w)
|
||||
currentUpdate.UpsertedWorkspaces = append(currentUpdate.UpsertedWorkspaces, w)
|
||||
}
|
||||
|
||||
// delete agents before deleting workspaces, since the agents have workspace ID references
|
||||
@@ -1014,17 +1126,22 @@ func (t *tunnelUpdater) handleUpdate(update *proto.WorkspaceUpdate) error {
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to parse workspace ID: %w", err)
|
||||
}
|
||||
err = t.deleteAgent(workspaceID, agentID)
|
||||
deletedAgent, err := t.deleteAgentLocked(workspaceID, agentID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to delete agent: %w", err)
|
||||
}
|
||||
currentUpdate.DeletedAgents = append(currentUpdate.DeletedAgents, deletedAgent)
|
||||
}
|
||||
for _, dw := range update.DeletedWorkspaces {
|
||||
workspaceID, err := uuid.FromBytes(dw.Id)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to parse workspace ID: %w", err)
|
||||
}
|
||||
t.deleteWorkspace(workspaceID)
|
||||
deletedWorkspace, err := t.deleteWorkspaceLocked(workspaceID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to delete workspace: %w", err)
|
||||
}
|
||||
currentUpdate.DeletedWorkspaces = append(currentUpdate.DeletedWorkspaces, deletedWorkspace)
|
||||
}
|
||||
|
||||
// upsert agents last, after all workspaces have been added and deleted, since agents reference
|
||||
@@ -1038,17 +1155,18 @@ func (t *tunnelUpdater) handleUpdate(update *proto.WorkspaceUpdate) error {
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to parse workspace ID: %w", err)
|
||||
}
|
||||
a := agent{name: ua.Name, id: agentID}
|
||||
err = t.upsertAgent(workspaceID, a)
|
||||
a := &Agent{Name: ua.Name, ID: agentID, WorkspaceID: workspaceID}
|
||||
err = t.upsertAgentLocked(workspaceID, a)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to upsert agent: %w", err)
|
||||
}
|
||||
currentUpdate.UpsertedAgents = append(currentUpdate.UpsertedAgents, a)
|
||||
}
|
||||
allAgents := t.allAgentIDs()
|
||||
allAgents := t.allAgentIDsLocked()
|
||||
t.coordCtrl.SyncDestinations(allAgents)
|
||||
dnsNames := t.updateDNSNamesLocked()
|
||||
if t.dnsHostsSetter != nil {
|
||||
t.logger.Debug(context.Background(), "updating dns hosts")
|
||||
dnsNames := t.allDNSNames()
|
||||
err := t.dnsHostsSetter.SetDNSHosts(dnsNames)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to set DNS hosts: %w", err)
|
||||
@@ -1056,41 +1174,60 @@ func (t *tunnelUpdater) handleUpdate(update *proto.WorkspaceUpdate) error {
|
||||
} else {
|
||||
t.logger.Debug(context.Background(), "skipping setting DNS names because we have no setter")
|
||||
}
|
||||
if t.updateHandler != nil {
|
||||
t.logger.Debug(context.Background(), "calling update handler")
|
||||
err := t.updateHandler.Update(currentUpdate.Clone())
|
||||
if err != nil {
|
||||
t.logger.Error(context.Background(), "failed to call update handler", slog.Error(err))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) upsertWorkspace(w workspace) {
|
||||
old, ok := t.workspaces[w.id]
|
||||
func (t *tunnelUpdater) upsertWorkspaceLocked(w *Workspace) *Workspace {
|
||||
old, ok := t.workspaces[w.ID]
|
||||
if !ok {
|
||||
t.workspaces[w.id] = &w
|
||||
return
|
||||
t.workspaces[w.ID] = w
|
||||
return w
|
||||
}
|
||||
old.name = w.name
|
||||
old.Name = w.Name
|
||||
old.Status = w.Status
|
||||
old.ownerUsername = w.ownerUsername
|
||||
return w
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) deleteWorkspace(id uuid.UUID) {
|
||||
func (t *tunnelUpdater) deleteWorkspaceLocked(id uuid.UUID) (*Workspace, error) {
|
||||
w, ok := t.workspaces[id]
|
||||
if !ok {
|
||||
return nil, xerrors.Errorf("workspace %s not found", id)
|
||||
}
|
||||
delete(t.workspaces, id)
|
||||
return w, nil
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) upsertAgent(workspaceID uuid.UUID, a agent) error {
|
||||
func (t *tunnelUpdater) upsertAgentLocked(workspaceID uuid.UUID, a *Agent) error {
|
||||
w, ok := t.workspaces[workspaceID]
|
||||
if !ok {
|
||||
return xerrors.Errorf("workspace %s not found", workspaceID)
|
||||
}
|
||||
w.agents[a.id] = a
|
||||
w.agents[a.ID] = a
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) deleteAgent(workspaceID, id uuid.UUID) error {
|
||||
func (t *tunnelUpdater) deleteAgentLocked(workspaceID, id uuid.UUID) (*Agent, error) {
|
||||
w, ok := t.workspaces[workspaceID]
|
||||
if !ok {
|
||||
return xerrors.Errorf("workspace %s not found", workspaceID)
|
||||
return nil, xerrors.Errorf("workspace %s not found", workspaceID)
|
||||
}
|
||||
a, ok := w.agents[id]
|
||||
if !ok {
|
||||
return nil, xerrors.Errorf("agent %s not found in workspace %s", id, workspaceID)
|
||||
}
|
||||
delete(w.agents, id)
|
||||
return nil
|
||||
return a, nil
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) allAgentIDs() []uuid.UUID {
|
||||
func (t *tunnelUpdater) allAgentIDsLocked() []uuid.UUID {
|
||||
out := make([]uuid.UUID, 0, len(t.workspaces))
|
||||
for _, w := range t.workspaces {
|
||||
for id := range w.agents {
|
||||
@@ -1100,41 +1237,54 @@ func (t *tunnelUpdater) allAgentIDs() []uuid.UUID {
|
||||
return out
|
||||
}
|
||||
|
||||
func (t *tunnelUpdater) allDNSNames() map[dnsname.FQDN][]netip.Addr {
|
||||
// updateDNSNamesLocked updates the DNS names for all workspaces in the tunnelUpdater.
|
||||
// t.Mutex must be held.
|
||||
func (t *tunnelUpdater) updateDNSNamesLocked() map[dnsname.FQDN][]netip.Addr {
|
||||
names := make(map[dnsname.FQDN][]netip.Addr)
|
||||
for _, w := range t.workspaces {
|
||||
err := w.addAllDNSNames(names, t.ownerUsername)
|
||||
err := w.updateDNSNames()
|
||||
if err != nil {
|
||||
// This should never happen in production, because converting the FQDN only fails
|
||||
// if names are too long, and we put strict length limits on agent, workspace, and user
|
||||
// names.
|
||||
t.logger.Critical(context.Background(),
|
||||
"failed to include DNS name(s)",
|
||||
slog.F("workspace_id", w.id),
|
||||
slog.F("workspace_id", w.ID),
|
||||
slog.Error(err))
|
||||
}
|
||||
for _, a := range w.agents {
|
||||
for name, addrs := range a.Hosts {
|
||||
names[name] = addrs
|
||||
}
|
||||
}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
type TunnelAllOption func(t *tunnelAllWorkspaceUpdatesController)
|
||||
type TunnelAllOption func(t *TunnelAllWorkspaceUpdatesController)
|
||||
|
||||
// WithDNS configures the tunnelAllWorkspaceUpdatesController to set DNS names for all workspaces
|
||||
// and agents it learns about.
|
||||
func WithDNS(d DNSHostsSetter, ownerUsername string) TunnelAllOption {
|
||||
return func(t *tunnelAllWorkspaceUpdatesController) {
|
||||
return func(t *TunnelAllWorkspaceUpdatesController) {
|
||||
t.dnsHostSetter = d
|
||||
t.ownerUsername = ownerUsername
|
||||
}
|
||||
}
|
||||
|
||||
func WithHandler(h UpdatesHandler) TunnelAllOption {
|
||||
return func(t *TunnelAllWorkspaceUpdatesController) {
|
||||
t.updateHandler = h
|
||||
}
|
||||
}
|
||||
|
||||
// NewTunnelAllWorkspaceUpdatesController creates a WorkspaceUpdatesController that creates tunnels
|
||||
// (via the TunnelSrcCoordController) to all agents received over the WorkspaceUpdates RPC. If a
|
||||
// DNSHostSetter is provided, it also programs DNS hosts based on the agent and workspace names.
|
||||
func NewTunnelAllWorkspaceUpdatesController(
|
||||
logger slog.Logger, c *TunnelSrcCoordController, opts ...TunnelAllOption,
|
||||
) WorkspaceUpdatesController {
|
||||
t := &tunnelAllWorkspaceUpdatesController{logger: logger, coordCtrl: c}
|
||||
) *TunnelAllWorkspaceUpdatesController {
|
||||
t := &TunnelAllWorkspaceUpdatesController{logger: logger, coordCtrl: c}
|
||||
for _, opt := range opts {
|
||||
opt(t)
|
||||
}
|
||||
|
||||
+176
-22
@@ -7,6 +7,7 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
@@ -1451,10 +1452,35 @@ func (f *fakeDNSSetter) SetDNSHosts(hosts map[dnsname.FQDN][]netip.Addr) error {
|
||||
}
|
||||
}
|
||||
|
||||
func newFakeUpdateHandler(ctx context.Context, t testing.TB) *fakeUpdateHandler {
|
||||
return &fakeUpdateHandler{
|
||||
ctx: ctx,
|
||||
t: t,
|
||||
ch: make(chan tailnet.WorkspaceUpdate),
|
||||
}
|
||||
}
|
||||
|
||||
type fakeUpdateHandler struct {
|
||||
ctx context.Context
|
||||
t testing.TB
|
||||
ch chan tailnet.WorkspaceUpdate
|
||||
}
|
||||
|
||||
func (f *fakeUpdateHandler) Update(wu tailnet.WorkspaceUpdate) error {
|
||||
f.t.Helper()
|
||||
select {
|
||||
case <-f.ctx.Done():
|
||||
return timeoutOnFakeErr
|
||||
case f.ch <- wu:
|
||||
// OK
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func setupConnectedAllWorkspaceUpdatesController(
|
||||
ctx context.Context, t testing.TB, logger slog.Logger, opts ...tailnet.TunnelAllOption,
|
||||
) (
|
||||
*fakeCoordinatorClient, *fakeWorkspaceUpdateClient,
|
||||
*fakeCoordinatorClient, *fakeWorkspaceUpdateClient, *tailnet.TunnelAllWorkspaceUpdatesController,
|
||||
) {
|
||||
fConn := &fakeCoordinatee{}
|
||||
tsc := tailnet.NewTunnelSrcCoordController(logger, fConn)
|
||||
@@ -1484,7 +1510,7 @@ func setupConnectedAllWorkspaceUpdatesController(
|
||||
err := testutil.RequireRecvCtx(ctx, t, updateCW.Wait())
|
||||
require.ErrorIs(t, err, io.EOF)
|
||||
})
|
||||
return coordC, updateC
|
||||
return coordC, updateC, uut
|
||||
}
|
||||
|
||||
func TestTunnelAllWorkspaceUpdatesController_Initial(t *testing.T) {
|
||||
@@ -1492,9 +1518,12 @@ func TestTunnelAllWorkspaceUpdatesController_Initial(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := testutil.Logger(t)
|
||||
|
||||
fUH := newFakeUpdateHandler(ctx, t)
|
||||
fDNS := newFakeDNSSetter(ctx, t)
|
||||
coordC, updateC := setupConnectedAllWorkspaceUpdatesController(ctx, t, logger,
|
||||
tailnet.WithDNS(fDNS, "testy"))
|
||||
coordC, updateC, updateCtrl := setupConnectedAllWorkspaceUpdatesController(ctx, t, logger,
|
||||
tailnet.WithDNS(fDNS, "testy"),
|
||||
tailnet.WithHandler(fUH),
|
||||
)
|
||||
|
||||
// Initial update contains 2 workspaces with 1 & 2 agents, respectively
|
||||
w1ID := testUUID(1)
|
||||
@@ -1528,19 +1557,71 @@ func TestTunnelAllWorkspaceUpdatesController_Initial(t *testing.T) {
|
||||
require.Contains(t, adds, w2a1ID)
|
||||
require.Contains(t, adds, w2a2ID)
|
||||
|
||||
ws1a1IP := netip.MustParseAddr("fd60:627a:a42b:0101::")
|
||||
w2a1IP := netip.MustParseAddr("fd60:627a:a42b:0201::")
|
||||
w2a2IP := netip.MustParseAddr("fd60:627a:a42b:0202::")
|
||||
|
||||
// Also triggers setting DNS hosts
|
||||
expectedDNS := map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a1.w1.me.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w2a1.w2.me.coder.": {netip.MustParseAddr("fd60:627a:a42b:0201::")},
|
||||
"w2a2.w2.me.coder.": {netip.MustParseAddr("fd60:627a:a42b:0202::")},
|
||||
"w1a1.w1.testy.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w2a1.w2.testy.coder.": {netip.MustParseAddr("fd60:627a:a42b:0201::")},
|
||||
"w2a2.w2.testy.coder.": {netip.MustParseAddr("fd60:627a:a42b:0202::")},
|
||||
"w1.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1a1.w1.me.coder.": {ws1a1IP},
|
||||
"w2a1.w2.me.coder.": {w2a1IP},
|
||||
"w2a2.w2.me.coder.": {w2a2IP},
|
||||
"w1a1.w1.testy.coder.": {ws1a1IP},
|
||||
"w2a1.w2.testy.coder.": {w2a1IP},
|
||||
"w2a2.w2.testy.coder.": {w2a2IP},
|
||||
"w1.coder.": {ws1a1IP},
|
||||
}
|
||||
dnsCall := testutil.RequireRecvCtx(ctx, t, fDNS.calls)
|
||||
require.Equal(t, expectedDNS, dnsCall.hosts)
|
||||
testutil.RequireSendCtx(ctx, t, dnsCall.err, nil)
|
||||
|
||||
currentState := tailnet.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*tailnet.Workspace{
|
||||
{ID: w1ID, Name: "w1"},
|
||||
{ID: w2ID, Name: "w2"},
|
||||
},
|
||||
UpsertedAgents: []*tailnet.Agent{
|
||||
{
|
||||
ID: w1a1ID, Name: "w1a1", WorkspaceID: w1ID,
|
||||
Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w1.coder.": {ws1a1IP},
|
||||
"w1a1.w1.me.coder.": {ws1a1IP},
|
||||
"w1a1.w1.testy.coder.": {ws1a1IP},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: w2a1ID, Name: "w2a1", WorkspaceID: w2ID,
|
||||
Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w2a1.w2.me.coder.": {w2a1IP},
|
||||
"w2a1.w2.testy.coder.": {w2a1IP},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: w2a2ID, Name: "w2a2", WorkspaceID: w2ID,
|
||||
Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w2a2.w2.me.coder.": {w2a2IP},
|
||||
"w2a2.w2.testy.coder.": {w2a2IP},
|
||||
},
|
||||
},
|
||||
},
|
||||
DeletedWorkspaces: []*tailnet.Workspace{},
|
||||
DeletedAgents: []*tailnet.Agent{},
|
||||
}
|
||||
|
||||
// And the callback
|
||||
cbUpdate := testutil.RequireRecvCtx(ctx, t, fUH.ch)
|
||||
require.Equal(t, currentState, cbUpdate)
|
||||
|
||||
// Current recvState should match
|
||||
recvState, err := updateCtrl.CurrentState()
|
||||
require.NoError(t, err)
|
||||
slices.SortFunc(recvState.UpsertedWorkspaces, func(a, b *tailnet.Workspace) int {
|
||||
return strings.Compare(a.Name, b.Name)
|
||||
})
|
||||
slices.SortFunc(recvState.UpsertedAgents, func(a, b *tailnet.Agent) int {
|
||||
return strings.Compare(a.Name, b.Name)
|
||||
})
|
||||
require.Equal(t, currentState, recvState)
|
||||
}
|
||||
|
||||
func TestTunnelAllWorkspaceUpdatesController_DeleteAgent(t *testing.T) {
|
||||
@@ -1548,13 +1629,19 @@ func TestTunnelAllWorkspaceUpdatesController_DeleteAgent(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := testutil.Logger(t)
|
||||
|
||||
fUH := newFakeUpdateHandler(ctx, t)
|
||||
fDNS := newFakeDNSSetter(ctx, t)
|
||||
coordC, updateC := setupConnectedAllWorkspaceUpdatesController(ctx, t, logger,
|
||||
tailnet.WithDNS(fDNS, "testy"))
|
||||
coordC, updateC, updateCtrl := setupConnectedAllWorkspaceUpdatesController(ctx, t, logger,
|
||||
tailnet.WithDNS(fDNS, "testy"),
|
||||
tailnet.WithHandler(fUH),
|
||||
)
|
||||
|
||||
w1ID := testUUID(1)
|
||||
w1a1ID := testUUID(1, 1)
|
||||
w1a2ID := testUUID(1, 2)
|
||||
ws1a1IP := netip.MustParseAddr("fd60:627a:a42b:0101::")
|
||||
ws1a2IP := netip.MustParseAddr("fd60:627a:a42b:0102::")
|
||||
|
||||
initUp := &proto.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*proto.Workspace{
|
||||
{Id: w1ID[:], Name: "w1"},
|
||||
@@ -1574,14 +1661,37 @@ func TestTunnelAllWorkspaceUpdatesController_DeleteAgent(t *testing.T) {
|
||||
|
||||
// DNS for w1a1
|
||||
expectedDNS := map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a1.w1.testy.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1a1.w1.me.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1a1.w1.testy.coder.": {ws1a1IP},
|
||||
"w1a1.w1.me.coder.": {ws1a1IP},
|
||||
"w1.coder.": {ws1a1IP},
|
||||
}
|
||||
dnsCall := testutil.RequireRecvCtx(ctx, t, fDNS.calls)
|
||||
require.Equal(t, expectedDNS, dnsCall.hosts)
|
||||
testutil.RequireSendCtx(ctx, t, dnsCall.err, nil)
|
||||
|
||||
initRecvUp := tailnet.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*tailnet.Workspace{
|
||||
{ID: w1ID, Name: "w1"},
|
||||
},
|
||||
UpsertedAgents: []*tailnet.Agent{
|
||||
{ID: w1a1ID, Name: "w1a1", WorkspaceID: w1ID, Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a1.w1.testy.coder.": {ws1a1IP},
|
||||
"w1a1.w1.me.coder.": {ws1a1IP},
|
||||
"w1.coder.": {ws1a1IP},
|
||||
}},
|
||||
},
|
||||
DeletedWorkspaces: []*tailnet.Workspace{},
|
||||
DeletedAgents: []*tailnet.Agent{},
|
||||
}
|
||||
|
||||
cbUpdate := testutil.RequireRecvCtx(ctx, t, fUH.ch)
|
||||
require.Equal(t, initRecvUp, cbUpdate)
|
||||
|
||||
// Current state should match initial
|
||||
state, err := updateCtrl.CurrentState()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, initRecvUp, state)
|
||||
|
||||
// Send update that removes w1a1 and adds w1a2
|
||||
agentUpdate := &proto.WorkspaceUpdate{
|
||||
UpsertedAgents: []*proto.Agent{
|
||||
@@ -1606,13 +1716,51 @@ func TestTunnelAllWorkspaceUpdatesController_DeleteAgent(t *testing.T) {
|
||||
|
||||
// DNS contains only w1a2
|
||||
expectedDNS = map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a2.w1.testy.coder.": {netip.MustParseAddr("fd60:627a:a42b:0102::")},
|
||||
"w1a2.w1.me.coder.": {netip.MustParseAddr("fd60:627a:a42b:0102::")},
|
||||
"w1.coder.": {netip.MustParseAddr("fd60:627a:a42b:0102::")},
|
||||
"w1a2.w1.testy.coder.": {ws1a2IP},
|
||||
"w1a2.w1.me.coder.": {ws1a2IP},
|
||||
"w1.coder.": {ws1a2IP},
|
||||
}
|
||||
dnsCall = testutil.RequireRecvCtx(ctx, t, fDNS.calls)
|
||||
require.Equal(t, expectedDNS, dnsCall.hosts)
|
||||
testutil.RequireSendCtx(ctx, t, dnsCall.err, nil)
|
||||
|
||||
cbUpdate = testutil.RequireRecvCtx(ctx, t, fUH.ch)
|
||||
sndRecvUpdate := tailnet.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*tailnet.Workspace{},
|
||||
UpsertedAgents: []*tailnet.Agent{
|
||||
{ID: w1a2ID, Name: "w1a2", WorkspaceID: w1ID, Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a2.w1.testy.coder.": {ws1a2IP},
|
||||
"w1a2.w1.me.coder.": {ws1a2IP},
|
||||
"w1.coder.": {ws1a2IP},
|
||||
}},
|
||||
},
|
||||
DeletedWorkspaces: []*tailnet.Workspace{},
|
||||
DeletedAgents: []*tailnet.Agent{
|
||||
{ID: w1a1ID, Name: "w1a1", WorkspaceID: w1ID, Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a1.w1.testy.coder.": {ws1a1IP},
|
||||
"w1a1.w1.me.coder.": {ws1a1IP},
|
||||
"w1.coder.": {ws1a1IP},
|
||||
}},
|
||||
},
|
||||
}
|
||||
require.Equal(t, sndRecvUpdate, cbUpdate)
|
||||
|
||||
state, err = updateCtrl.CurrentState()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tailnet.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*tailnet.Workspace{
|
||||
{ID: w1ID, Name: "w1"},
|
||||
},
|
||||
UpsertedAgents: []*tailnet.Agent{
|
||||
{ID: w1a2ID, Name: "w1a2", WorkspaceID: w1ID, Hosts: map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a2.w1.testy.coder.": {ws1a2IP},
|
||||
"w1a2.w1.me.coder.": {ws1a2IP},
|
||||
"w1.coder.": {ws1a2IP},
|
||||
}},
|
||||
},
|
||||
DeletedWorkspaces: []*tailnet.Workspace{},
|
||||
DeletedAgents: []*tailnet.Agent{},
|
||||
}, state)
|
||||
}
|
||||
|
||||
func TestTunnelAllWorkspaceUpdatesController_DNSError(t *testing.T) {
|
||||
@@ -1635,6 +1783,8 @@ func TestTunnelAllWorkspaceUpdatesController_DNSError(t *testing.T) {
|
||||
|
||||
w1ID := testUUID(1)
|
||||
w1a1ID := testUUID(1, 1)
|
||||
ws1a1IP := netip.MustParseAddr("fd60:627a:a42b:0101::")
|
||||
|
||||
initUp := &proto.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*proto.Workspace{
|
||||
{Id: w1ID[:], Name: "w1"},
|
||||
@@ -1648,9 +1798,9 @@ func TestTunnelAllWorkspaceUpdatesController_DNSError(t *testing.T) {
|
||||
|
||||
// DNS for w1a1
|
||||
expectedDNS := map[dnsname.FQDN][]netip.Addr{
|
||||
"w1a1.w1.me.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1a1.w1.testy.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1.coder.": {netip.MustParseAddr("fd60:627a:a42b:0101::")},
|
||||
"w1a1.w1.me.coder.": {ws1a1IP},
|
||||
"w1a1.w1.testy.coder.": {ws1a1IP},
|
||||
"w1.coder.": {ws1a1IP},
|
||||
}
|
||||
dnsCall := testutil.RequireRecvCtx(ctx, t, fDNS.calls)
|
||||
require.Equal(t, expectedDNS, dnsCall.hosts)
|
||||
@@ -1778,6 +1928,10 @@ type fakeWorkspaceUpdatesController struct {
|
||||
calls chan *newWorkspaceUpdatesCall
|
||||
}
|
||||
|
||||
func (*fakeWorkspaceUpdatesController) CurrentState() *proto.WorkspaceUpdate {
|
||||
panic("unimplemented")
|
||||
}
|
||||
|
||||
type newWorkspaceUpdatesCall struct {
|
||||
client tailnet.WorkspaceUpdatesClient
|
||||
resp chan<- tailnet.CloserWaiter
|
||||
|
||||
Reference in New Issue
Block a user