mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add support for WorkspaceUpdates to WebsocketDialer (#15534)
closes #14730 Adds support for WorkspaceUpdates to the WebsocketDialer. This allows us to dial the new endpoint added in #14847 and connect it up to a `tailnet.Controllers` to connect to all agents over the tailnet. I refactored the fakeWorkspaceUpdatesProvider to a mock and moved it to `tailnettest` so it could be more easily reused. The Mock is a little more full-featured.
This commit is contained in:
@@ -25,14 +25,26 @@ var permanentErrorStatuses = []int{
|
||||
}
|
||||
|
||||
type WebsocketDialer struct {
|
||||
logger slog.Logger
|
||||
dialOptions *websocket.DialOptions
|
||||
url *url.URL
|
||||
logger slog.Logger
|
||||
dialOptions *websocket.DialOptions
|
||||
url *url.URL
|
||||
// workspaceUpdatesReq != nil means that the dialer should call the WorkspaceUpdates RPC and
|
||||
// return the corresponding client
|
||||
workspaceUpdatesReq *proto.WorkspaceUpdatesRequest
|
||||
|
||||
resumeTokenFailed bool
|
||||
connected chan error
|
||||
isFirst bool
|
||||
}
|
||||
|
||||
type WebsocketDialerOption func(*WebsocketDialer)
|
||||
|
||||
func WithWorkspaceUpdates(req *proto.WorkspaceUpdatesRequest) WebsocketDialerOption {
|
||||
return func(w *WebsocketDialer) {
|
||||
w.workspaceUpdatesReq = req
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebsocketDialer) Dial(ctx context.Context, r tailnet.ResumeTokenController,
|
||||
) (
|
||||
tailnet.ControlProtocolClients, error,
|
||||
@@ -41,14 +53,27 @@ func (w *WebsocketDialer) Dial(ctx context.Context, r tailnet.ResumeTokenControl
|
||||
|
||||
u := new(url.URL)
|
||||
*u = *w.url
|
||||
q := u.Query()
|
||||
if r != nil && !w.resumeTokenFailed {
|
||||
if token, ok := r.Token(); ok {
|
||||
q := u.Query()
|
||||
q.Set("resume_token", token)
|
||||
u.RawQuery = q.Encode()
|
||||
w.logger.Debug(ctx, "using resume token on dial")
|
||||
}
|
||||
}
|
||||
// The current version includes additions
|
||||
//
|
||||
// 2.1 GetAnnouncementBanners on the Agent API (version locked to Tailnet API)
|
||||
// 2.2 PostTelemetry on the Tailnet API
|
||||
// 2.3 RefreshResumeToken, WorkspaceUpdates
|
||||
//
|
||||
// Resume tokens and telemetry are optional, and fail gracefully. So we use version 2.0 for
|
||||
// maximum compatibility if we don't need WorkspaceUpdates. If we do, we use 2.3.
|
||||
if w.workspaceUpdatesReq != nil {
|
||||
q.Add("version", "2.3")
|
||||
} else {
|
||||
q.Add("version", "2.0")
|
||||
}
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
// nolint:bodyclose
|
||||
ws, res, err := websocket.Dial(ctx, u.String(), w.dialOptions)
|
||||
@@ -115,12 +140,23 @@ func (w *WebsocketDialer) Dial(ctx context.Context, r tailnet.ResumeTokenControl
|
||||
return tailnet.ControlProtocolClients{}, err
|
||||
}
|
||||
|
||||
var updates tailnet.WorkspaceUpdatesClient
|
||||
if w.workspaceUpdatesReq != nil {
|
||||
updates, err = client.WorkspaceUpdates(context.Background(), w.workspaceUpdatesReq)
|
||||
if err != nil {
|
||||
w.logger.Debug(ctx, "failed to create WorkspaceUpdates stream", slog.Error(err))
|
||||
_ = ws.Close(websocket.StatusInternalError, "")
|
||||
return tailnet.ControlProtocolClients{}, err
|
||||
}
|
||||
}
|
||||
|
||||
return tailnet.ControlProtocolClients{
|
||||
Closer: client.DRPCConn(),
|
||||
Coordinator: coord,
|
||||
DERP: derps,
|
||||
ResumeToken: client,
|
||||
Telemetry: client,
|
||||
Closer: client.DRPCConn(),
|
||||
Coordinator: coord,
|
||||
DERP: derps,
|
||||
ResumeToken: client,
|
||||
Telemetry: client,
|
||||
WorkspaceUpdates: updates,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -128,12 +164,19 @@ func (w *WebsocketDialer) Connected() <-chan error {
|
||||
return w.connected
|
||||
}
|
||||
|
||||
func NewWebsocketDialer(logger slog.Logger, u *url.URL, opts *websocket.DialOptions) *WebsocketDialer {
|
||||
return &WebsocketDialer{
|
||||
func NewWebsocketDialer(
|
||||
logger slog.Logger, u *url.URL, websocketOptions *websocket.DialOptions,
|
||||
dialerOptions ...WebsocketDialerOption,
|
||||
) *WebsocketDialer {
|
||||
w := &WebsocketDialer{
|
||||
logger: logger,
|
||||
dialOptions: opts,
|
||||
dialOptions: websocketOptions,
|
||||
url: u,
|
||||
connected: make(chan error, 1),
|
||||
isFirst: true,
|
||||
}
|
||||
for _, o := range dialerOptions {
|
||||
o(w)
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
@@ -9,8 +9,10 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"nhooyr.io/websocket"
|
||||
"tailscale.com/tailcfg"
|
||||
|
||||
@@ -21,7 +23,7 @@ import (
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
tailnetproto "github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/coder/v2/tailnet/tailnettest"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
@@ -102,6 +104,7 @@ func TestWebsocketDialer_TokenController(t *testing.T) {
|
||||
require.Equal(t, "", gotToken)
|
||||
|
||||
clients = testutil.RequireRecvCtx(ctx, t, clientCh)
|
||||
require.Nil(t, clients.WorkspaceUpdates)
|
||||
clients.Closer.Close()
|
||||
|
||||
err = testutil.RequireRecvCtx(ctx, t, wsErr)
|
||||
@@ -273,7 +276,7 @@ func TestWebsocketDialer_UplevelVersion(t *testing.T) {
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
|
||||
|
||||
svr := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
sVer := apiversion.New(proto.CurrentMajor, proto.CurrentMinor-1)
|
||||
sVer := apiversion.New(2, 2)
|
||||
|
||||
// the following matches what Coderd does;
|
||||
// c.f. coderd/workspaceagents.go: workspaceAgentClientCoordinate
|
||||
@@ -291,7 +294,10 @@ func TestWebsocketDialer_UplevelVersion(t *testing.T) {
|
||||
svrURL, err := url.Parse(svr.URL)
|
||||
require.NoError(t, err)
|
||||
|
||||
uut := workspacesdk.NewWebsocketDialer(logger, svrURL, &websocket.DialOptions{})
|
||||
uut := workspacesdk.NewWebsocketDialer(
|
||||
logger, svrURL, &websocket.DialOptions{},
|
||||
workspacesdk.WithWorkspaceUpdates(&tailnetproto.WorkspaceUpdatesRequest{}),
|
||||
)
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
@@ -307,6 +313,84 @@ func TestWebsocketDialer_UplevelVersion(t *testing.T) {
|
||||
require.NotEmpty(t, sdkErr.Helper)
|
||||
}
|
||||
|
||||
func TestWebsocketDialer_WorkspaceUpdates(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := slogtest.Make(t, &slogtest.Options{
|
||||
IgnoreErrors: true,
|
||||
}).Leveled(slog.LevelDebug)
|
||||
|
||||
fCoord := tailnettest.NewFakeCoordinator()
|
||||
var coord tailnet.Coordinator = fCoord
|
||||
coordPtr := atomic.Pointer[tailnet.Coordinator]{}
|
||||
coordPtr.Store(&coord)
|
||||
ctrl := gomock.NewController(t)
|
||||
mProvider := tailnettest.NewMockWorkspaceUpdatesProvider(ctrl)
|
||||
|
||||
svc, err := tailnet.NewClientService(tailnet.ClientServiceOptions{
|
||||
Logger: logger,
|
||||
CoordPtr: &coordPtr,
|
||||
DERPMapUpdateFrequency: time.Hour,
|
||||
DERPMapFn: func() *tailcfg.DERPMap { return &tailcfg.DERPMap{} },
|
||||
WorkspaceUpdatesProvider: mProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
wsErr := make(chan error, 1)
|
||||
svr := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// need 2.3 for WorkspaceUpdates RPC
|
||||
cVer := r.URL.Query().Get("version")
|
||||
assert.Equal(t, "2.3", cVer)
|
||||
|
||||
sws, err := websocket.Accept(w, r, nil)
|
||||
if !assert.NoError(t, err) {
|
||||
return
|
||||
}
|
||||
wsCtx, nc := codersdk.WebsocketNetConn(ctx, sws, websocket.MessageBinary)
|
||||
// streamID can be empty because we don't call RPCs in this test.
|
||||
wsErr <- svc.ServeConnV2(wsCtx, nc, tailnet.StreamID{})
|
||||
}))
|
||||
defer svr.Close()
|
||||
svrURL, err := url.Parse(svr.URL)
|
||||
require.NoError(t, err)
|
||||
|
||||
userID := uuid.UUID{88}
|
||||
|
||||
mSub := tailnettest.NewMockSubscription(ctrl)
|
||||
updateCh := make(chan *tailnetproto.WorkspaceUpdate, 1)
|
||||
mProvider.EXPECT().Subscribe(gomock.Any(), userID).Times(1).Return(mSub, nil)
|
||||
mSub.EXPECT().Updates().MinTimes(1).Return(updateCh)
|
||||
mSub.EXPECT().Close().Times(1).Return(nil)
|
||||
|
||||
uut := workspacesdk.NewWebsocketDialer(
|
||||
logger, svrURL, &websocket.DialOptions{},
|
||||
workspacesdk.WithWorkspaceUpdates(&tailnetproto.WorkspaceUpdatesRequest{
|
||||
WorkspaceOwnerId: userID[:],
|
||||
}),
|
||||
)
|
||||
|
||||
clients, err := uut.Dial(ctx, nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, clients.WorkspaceUpdates)
|
||||
|
||||
wsID := uuid.UUID{99}
|
||||
expectedUpdate := &tailnetproto.WorkspaceUpdate{
|
||||
UpsertedWorkspaces: []*tailnetproto.Workspace{
|
||||
{Id: wsID[:]},
|
||||
},
|
||||
}
|
||||
updateCh <- expectedUpdate
|
||||
|
||||
gotUpdate, err := clients.WorkspaceUpdates.Recv()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, wsID[:], gotUpdate.GetUpsertedWorkspaces()[0].GetId())
|
||||
|
||||
clients.Closer.Close()
|
||||
|
||||
err = testutil.RequireRecvCtx(ctx, t, wsErr)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
type fakeResumeTokenController struct {
|
||||
ctx context.Context
|
||||
t testing.TB
|
||||
|
||||
@@ -216,17 +216,6 @@ func (c *Client) DialAgent(dialCtx context.Context, agentID uuid.UUID, options *
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse url: %w", err)
|
||||
}
|
||||
q := coordinateURL.Query()
|
||||
// The current version includes additions
|
||||
//
|
||||
// 2.1 GetAnnouncementBanners on the Agent API (version locked to Tailnet API)
|
||||
// 2.2 PostTelemetry on the Tailnet API
|
||||
// 2.3 RefreshResumeToken, WorkspaceUpdates
|
||||
//
|
||||
// Since resume tokens and telemetry are optional, and fail gracefully, and we don't use
|
||||
// WorkspaceUpdates to talk to a single agent, we ask for version 2.0 for maximum compatibility
|
||||
q.Add("version", "2.0")
|
||||
coordinateURL.RawQuery = q.Encode()
|
||||
|
||||
dialer := NewWebsocketDialer(options.Logger, coordinateURL, &websocket.DialOptions{
|
||||
HTTPClient: c.client.HTTPClient,
|
||||
|
||||
Reference in New Issue
Block a user