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:
@@ -158,7 +158,7 @@ func New(ctx context.Context, opts *Options) (*Server, error) {
|
||||
// TODO: Probably do some version checking here
|
||||
info, err := client.SDKClient.BuildInfo(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildinfo: %w", errors.Join(
|
||||
return nil, xerrors.Errorf("buildinfo: %w", errors.Join(
|
||||
xerrors.Errorf("unable to fetch build info from primary coderd. Are you sure %q is a coderd instance?", opts.DashboardURL),
|
||||
err,
|
||||
))
|
||||
|
||||
@@ -5,7 +5,6 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
@@ -23,6 +22,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/workspaceapps"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
agpl "github.com/coder/coder/v2/tailnet"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
)
|
||||
|
||||
// Client is a HTTP client for a subset of Coder API routes that external
|
||||
@@ -438,6 +438,9 @@ func (c *Client) DialCoordinator(ctx context.Context) (agpl.MultiAgentConn, erro
|
||||
cancel()
|
||||
return nil, xerrors.Errorf("parse url: %w", err)
|
||||
}
|
||||
q := coordinateURL.Query()
|
||||
q.Add("version", agpl.CurrentVersion.String())
|
||||
coordinateURL.RawQuery = q.Encode()
|
||||
coordinateHeaders := make(http.Header)
|
||||
tokenHeader := codersdk.SessionTokenHeader
|
||||
if c.SDKClient.SessionTokenHeader != "" {
|
||||
@@ -457,10 +460,24 @@ func (c *Client) DialCoordinator(ctx context.Context) (agpl.MultiAgentConn, erro
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, conn)
|
||||
|
||||
nc := websocket.NetConn(ctx, conn, websocket.MessageText)
|
||||
nc := websocket.NetConn(ctx, conn, websocket.MessageBinary)
|
||||
client, err := agpl.NewDRPCClient(nc)
|
||||
if err != nil {
|
||||
logger.Debug(ctx, "failed to create DRPCClient", slog.Error(err))
|
||||
_ = conn.Close(websocket.StatusInternalError, "")
|
||||
return nil, xerrors.Errorf("failed to create DRPCClient: %w", err)
|
||||
}
|
||||
protocol, err := client.Coordinate(ctx)
|
||||
if err != nil {
|
||||
logger.Debug(ctx, "failed to reach the Coordinate endpoint", slog.Error(err))
|
||||
_ = conn.Close(websocket.StatusInternalError, "")
|
||||
return nil, xerrors.Errorf("failed to reach the Coordinate endpoint: %w", err)
|
||||
}
|
||||
|
||||
rma := remoteMultiAgentHandler{
|
||||
sdk: c,
|
||||
nc: nc,
|
||||
logger: logger,
|
||||
protocol: protocol,
|
||||
cancel: cancel,
|
||||
legacyAgentCache: map[uuid.UUID]bool{},
|
||||
}
|
||||
@@ -471,103 +488,75 @@ func (c *Client) DialCoordinator(ctx context.Context) (agpl.MultiAgentConn, erro
|
||||
OnSubscribe: rma.OnSubscribe,
|
||||
OnUnsubscribe: rma.OnUnsubscribe,
|
||||
OnNodeUpdate: rma.OnNodeUpdate,
|
||||
OnRemove: func(agpl.Queue) { conn.Close(websocket.StatusGoingAway, "closed") },
|
||||
OnRemove: rma.OnRemove,
|
||||
}).Init()
|
||||
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
ma.Close()
|
||||
_ = conn.Close(websocket.StatusGoingAway, "closed")
|
||||
}()
|
||||
|
||||
go func() {
|
||||
defer cancel()
|
||||
dec := json.NewDecoder(nc)
|
||||
for {
|
||||
var msg CoordinateNodes
|
||||
err := dec.Decode(&msg)
|
||||
if err != nil {
|
||||
if xerrors.Is(err, io.EOF) {
|
||||
logger.Info(ctx, "websocket connection severed", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
logger.Error(ctx, "decode coordinator nodes", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
err = ma.Enqueue(msg.Nodes)
|
||||
if err != nil {
|
||||
logger.Error(ctx, "enqueue nodes from coordinator", slog.Error(err))
|
||||
continue
|
||||
}
|
||||
}
|
||||
}()
|
||||
rma.ma = ma
|
||||
go rma.respLoop()
|
||||
|
||||
return ma, nil
|
||||
}
|
||||
|
||||
type remoteMultiAgentHandler struct {
|
||||
sdk *Client
|
||||
nc net.Conn
|
||||
cancel func()
|
||||
sdk *Client
|
||||
logger slog.Logger
|
||||
protocol proto.DRPCTailnet_CoordinateClient
|
||||
ma *agpl.MultiAgent
|
||||
cancel func()
|
||||
|
||||
legacyMu sync.RWMutex
|
||||
legacyAgentCache map[uuid.UUID]bool
|
||||
legacySingleflight singleflight.Group[uuid.UUID, AgentIsLegacyResponse]
|
||||
}
|
||||
|
||||
func (a *remoteMultiAgentHandler) writeJSON(v interface{}) error {
|
||||
data, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("json marshal message: %w", err)
|
||||
}
|
||||
func (a *remoteMultiAgentHandler) respLoop() {
|
||||
{
|
||||
defer a.cancel()
|
||||
for {
|
||||
resp, err := a.protocol.Recv()
|
||||
if err != nil {
|
||||
if xerrors.Is(err, io.EOF) {
|
||||
a.logger.Info(context.Background(), "remote multiagent connection severed", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
// Set a deadline so that hung connections don't put back pressure on the system.
|
||||
// Node updates are tiny, so even the dinkiest connection can handle them if it's not hung.
|
||||
err = a.nc.SetWriteDeadline(time.Now().Add(agpl.WriteTimeout))
|
||||
if err != nil {
|
||||
a.cancel()
|
||||
return xerrors.Errorf("set write deadline: %w", err)
|
||||
}
|
||||
_, err = a.nc.Write(data)
|
||||
if err != nil {
|
||||
a.cancel()
|
||||
return xerrors.Errorf("write message: %w", err)
|
||||
}
|
||||
a.logger.Error(context.Background(), "error receiving multiagent responses", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
// nhooyr.io/websocket has a bugged implementation of deadlines on a websocket net.Conn. What they are
|
||||
// *supposed* to do is set a deadline for any subsequent writes to complete, otherwise the call to Write()
|
||||
// fails. What nhooyr.io/websocket does is set a timer, after which it expires the websocket write context.
|
||||
// If this timer fires, then the next write will fail *even if we set a new write deadline*. So, after
|
||||
// our successful write, it is important that we reset the deadline before it fires.
|
||||
err = a.nc.SetWriteDeadline(time.Time{})
|
||||
if err != nil {
|
||||
a.cancel()
|
||||
return xerrors.Errorf("clear write deadline: %w", err)
|
||||
err = a.ma.Enqueue(resp)
|
||||
if err != nil {
|
||||
a.logger.Error(context.Background(), "enqueue response from coordinator", slog.Error(err))
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *remoteMultiAgentHandler) OnNodeUpdate(_ uuid.UUID, node *agpl.Node) error {
|
||||
return a.writeJSON(CoordinateMessage{
|
||||
Type: CoordinateMessageTypeNodeUpdate,
|
||||
Node: node,
|
||||
})
|
||||
func (a *remoteMultiAgentHandler) OnNodeUpdate(_ uuid.UUID, node *proto.Node) error {
|
||||
return a.protocol.Send(&proto.CoordinateRequest{UpdateSelf: &proto.CoordinateRequest_UpdateSelf{Node: node}})
|
||||
}
|
||||
|
||||
func (a *remoteMultiAgentHandler) OnSubscribe(_ agpl.Queue, agentID uuid.UUID) (*agpl.Node, error) {
|
||||
return nil, a.writeJSON(CoordinateMessage{
|
||||
Type: CoordinateMessageTypeSubscribe,
|
||||
AgentID: agentID,
|
||||
})
|
||||
func (a *remoteMultiAgentHandler) OnSubscribe(_ agpl.Queue, agentID uuid.UUID) error {
|
||||
return a.protocol.Send(&proto.CoordinateRequest{AddTunnel: &proto.CoordinateRequest_Tunnel{Id: agentID[:]}})
|
||||
}
|
||||
|
||||
func (a *remoteMultiAgentHandler) OnUnsubscribe(_ agpl.Queue, agentID uuid.UUID) error {
|
||||
return a.writeJSON(CoordinateMessage{
|
||||
Type: CoordinateMessageTypeUnsubscribe,
|
||||
AgentID: agentID,
|
||||
})
|
||||
return a.protocol.Send(&proto.CoordinateRequest{RemoveTunnel: &proto.CoordinateRequest_Tunnel{Id: agentID[:]}})
|
||||
}
|
||||
|
||||
func (a *remoteMultiAgentHandler) OnRemove(_ agpl.Queue) {
|
||||
err := a.protocol.Send(&proto.CoordinateRequest{Disconnect: &proto.CoordinateRequest_Disconnect{}})
|
||||
if err != nil {
|
||||
a.logger.Warn(context.Background(), "failed to gracefully disconnect", slog.Error(err))
|
||||
}
|
||||
_ = a.protocol.CloseSend()
|
||||
}
|
||||
|
||||
func (a *remoteMultiAgentHandler) AgentIsLegacy(agentID uuid.UUID) bool {
|
||||
|
||||
@@ -18,8 +18,9 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/xerrors"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
"nhooyr.io/websocket"
|
||||
"tailscale.com/tailcfg"
|
||||
"tailscale.com/types/key"
|
||||
|
||||
"cdr.dev/slog"
|
||||
@@ -30,6 +31,7 @@ import (
|
||||
"github.com/coder/coder/v2/enterprise/tailnet"
|
||||
"github.com/coder/coder/v2/enterprise/wsproxy/wsproxysdk"
|
||||
agpl "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"
|
||||
)
|
||||
@@ -156,25 +158,48 @@ func TestDialCoordinator(t *testing.T) {
|
||||
t.Run("OK", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
var (
|
||||
ctx, cancel = context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
logger = slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
agentID = uuid.New()
|
||||
serverMultiAgent = tailnettest.NewMockMultiAgentConn(gomock.NewController(t))
|
||||
r = chi.NewRouter()
|
||||
srv = httptest.NewServer(r)
|
||||
ctx, cancel = context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
logger = slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
agentID = uuid.UUID{33}
|
||||
proxyID = uuid.UUID{44}
|
||||
mCoord = tailnettest.NewMockCoordinator(gomock.NewController(t))
|
||||
coord agpl.Coordinator = mCoord
|
||||
r = chi.NewRouter()
|
||||
srv = httptest.NewServer(r)
|
||||
)
|
||||
defer cancel()
|
||||
defer srv.Close()
|
||||
|
||||
coordPtr := atomic.Pointer[agpl.Coordinator]{}
|
||||
coordPtr.Store(&coord)
|
||||
cSrv, err := tailnet.NewClientService(
|
||||
logger, &coordPtr,
|
||||
time.Hour,
|
||||
func() *tailcfg.DERPMap { panic("not implemented") },
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
// buffer the channels here, so we don't need to read and write in goroutines to
|
||||
// avoid blocking
|
||||
reqs := make(chan *proto.CoordinateRequest, 100)
|
||||
resps := make(chan *proto.CoordinateResponse, 100)
|
||||
mCoord.EXPECT().Coordinate(gomock.Any(), proxyID, gomock.Any(), agpl.SingleTailnetTunnelAuth{}).
|
||||
Times(1).
|
||||
Return(reqs, resps)
|
||||
|
||||
serveMACErr := make(chan error, 1)
|
||||
r.Get("/api/v2/workspaceproxies/me/coordinate", func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := websocket.Accept(w, r, nil)
|
||||
require.NoError(t, err)
|
||||
nc := websocket.NetConn(r.Context(), conn, websocket.MessageText)
|
||||
defer serverMultiAgent.Close()
|
||||
|
||||
err = tailnet.ServeWorkspaceProxy(ctx, nc, serverMultiAgent)
|
||||
if !xerrors.Is(err, io.EOF) {
|
||||
assert.NoError(t, err)
|
||||
if !assert.NoError(t, err) {
|
||||
return
|
||||
}
|
||||
version := r.URL.Query().Get("version")
|
||||
if !assert.Equal(t, version, agpl.CurrentVersion.String()) {
|
||||
return
|
||||
}
|
||||
nc := websocket.NetConn(r.Context(), conn, websocket.MessageBinary)
|
||||
err = cSrv.ServeMultiAgentClient(ctx, version, nc, proxyID)
|
||||
serveMACErr <- err
|
||||
})
|
||||
r.Get("/api/v2/workspaceagents/{workspaceagent}/legacy", func(w http.ResponseWriter, r *http.Request) {
|
||||
httpapi.Write(ctx, w, http.StatusOK, wsproxysdk.AgentIsLegacyResponse{
|
||||
@@ -188,51 +213,50 @@ func TestDialCoordinator(t *testing.T) {
|
||||
client := wsproxysdk.New(u)
|
||||
client.SDKClient.SetLogger(logger)
|
||||
|
||||
expected := []*agpl.Node{{
|
||||
ID: 55,
|
||||
AsOf: time.Unix(1689653252, 0),
|
||||
Key: key.NewNode().Public(),
|
||||
DiscoKey: key.NewDisco().Public(),
|
||||
PreferredDERP: 0,
|
||||
DERPLatency: map[string]float64{
|
||||
"0": 1.0,
|
||||
peerID := uuid.UUID{55}
|
||||
peerNodeKey, err := key.NewNode().Public().MarshalBinary()
|
||||
require.NoError(t, err)
|
||||
peerDiscoKey, err := key.NewDisco().Public().MarshalText()
|
||||
require.NoError(t, err)
|
||||
expected := &proto.CoordinateResponse{PeerUpdates: []*proto.CoordinateResponse_PeerUpdate{{
|
||||
Id: peerID[:],
|
||||
Node: &proto.Node{
|
||||
Id: 55,
|
||||
AsOf: timestamppb.New(time.Unix(1689653252, 0)),
|
||||
Key: peerNodeKey[:],
|
||||
Disco: string(peerDiscoKey),
|
||||
PreferredDerp: 0,
|
||||
DerpLatency: map[string]float64{
|
||||
"0": 1.0,
|
||||
},
|
||||
DerpForcedWebsocket: map[int32]string{},
|
||||
Addresses: []string{netip.PrefixFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4}), 128).String()},
|
||||
AllowedIps: []string{netip.PrefixFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4}), 128).String()},
|
||||
Endpoints: []string{"192.168.1.1:18842"},
|
||||
},
|
||||
DERPForcedWebsocket: map[int]string{},
|
||||
Addresses: []netip.Prefix{netip.PrefixFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4}), 128)},
|
||||
AllowedIPs: []netip.Prefix{netip.PrefixFrom(netip.AddrFrom16([16]byte{1, 2, 3, 4}), 128)},
|
||||
Endpoints: []string{"192.168.1.1:18842"},
|
||||
}}
|
||||
sendNode := make(chan struct{})
|
||||
|
||||
serverMultiAgent.EXPECT().NextUpdate(gomock.Any()).AnyTimes().
|
||||
DoAndReturn(func(ctx context.Context) ([]*agpl.Node, bool) {
|
||||
select {
|
||||
case <-sendNode:
|
||||
return expected, true
|
||||
case <-ctx.Done():
|
||||
return nil, false
|
||||
}
|
||||
})
|
||||
}}}
|
||||
|
||||
rma, err := client.DialCoordinator(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Subscribe
|
||||
{
|
||||
ch := make(chan struct{})
|
||||
serverMultiAgent.EXPECT().SubscribeAgent(agentID).Do(func(uuid.UUID) {
|
||||
close(ch)
|
||||
})
|
||||
require.NoError(t, rma.SubscribeAgent(agentID))
|
||||
waitOrCancel(ctx, t, ch)
|
||||
|
||||
req := testutil.RequireRecvCtx(ctx, t, reqs)
|
||||
require.Equal(t, agentID[:], req.GetAddTunnel().GetId())
|
||||
}
|
||||
// Read updated agent node
|
||||
{
|
||||
sendNode <- struct{}{}
|
||||
got, ok := rma.NextUpdate(ctx)
|
||||
resps <- expected
|
||||
|
||||
resp, ok := rma.NextUpdate(ctx)
|
||||
assert.True(t, ok)
|
||||
got[0].AsOf = got[0].AsOf.In(time.Local)
|
||||
assert.Equal(t, *expected[0], *got[0])
|
||||
updates := resp.GetPeerUpdates()
|
||||
assert.Len(t, updates, 1)
|
||||
eq, err := updates[0].GetNode().Equal(expected.GetPeerUpdates()[0].GetNode())
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, eq)
|
||||
}
|
||||
// Check legacy
|
||||
{
|
||||
@@ -241,45 +265,38 @@ func TestDialCoordinator(t *testing.T) {
|
||||
}
|
||||
// UpdateSelf
|
||||
{
|
||||
ch := make(chan struct{})
|
||||
serverMultiAgent.EXPECT().UpdateSelf(gomock.Any()).Do(func(node *agpl.Node) {
|
||||
node.AsOf = node.AsOf.In(time.Local)
|
||||
assert.Equal(t, expected[0], node)
|
||||
close(ch)
|
||||
})
|
||||
require.NoError(t, rma.UpdateSelf(expected[0]))
|
||||
waitOrCancel(ctx, t, ch)
|
||||
require.NoError(t, rma.UpdateSelf(expected.PeerUpdates[0].GetNode()))
|
||||
|
||||
req := testutil.RequireRecvCtx(ctx, t, reqs)
|
||||
eq, err := req.GetUpdateSelf().GetNode().Equal(expected.PeerUpdates[0].GetNode())
|
||||
require.NoError(t, err)
|
||||
require.True(t, eq)
|
||||
}
|
||||
// Unsubscribe
|
||||
{
|
||||
ch := make(chan struct{})
|
||||
serverMultiAgent.EXPECT().UnsubscribeAgent(agentID).Do(func(uuid.UUID) {
|
||||
close(ch)
|
||||
})
|
||||
require.NoError(t, rma.UnsubscribeAgent(agentID))
|
||||
waitOrCancel(ctx, t, ch)
|
||||
|
||||
req := testutil.RequireRecvCtx(ctx, t, reqs)
|
||||
require.Equal(t, agentID[:], req.GetRemoveTunnel().GetId())
|
||||
}
|
||||
// Close
|
||||
{
|
||||
ch := make(chan struct{})
|
||||
serverMultiAgent.EXPECT().Close().Do(func() {
|
||||
close(ch)
|
||||
})
|
||||
require.NoError(t, rma.Close())
|
||||
waitOrCancel(ctx, t, ch)
|
||||
|
||||
req := testutil.RequireRecvCtx(ctx, t, reqs)
|
||||
require.NotNil(t, req.Disconnect)
|
||||
close(resps)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timeout waiting for req close")
|
||||
case _, ok := <-reqs:
|
||||
require.False(t, ok, "didn't close requests")
|
||||
}
|
||||
require.Error(t, testutil.RequireRecvCtx(ctx, t, serveMACErr))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func waitOrCancel(ctx context.Context, t testing.TB, ch <-chan struct{}) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-ch:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for channel")
|
||||
}
|
||||
}
|
||||
|
||||
type ResponseRecorder struct {
|
||||
rw *httptest.ResponseRecorder
|
||||
wasWritten atomic.Bool
|
||||
|
||||
Reference in New Issue
Block a user