mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add resume support to coordinator connections (#14234)
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -24,6 +25,7 @@ import (
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/quartz"
|
||||
"github.com/coder/retry"
|
||||
)
|
||||
|
||||
@@ -61,6 +63,7 @@ type tailnetAPIConnector struct {
|
||||
|
||||
agentID uuid.UUID
|
||||
coordinateURL string
|
||||
clock quartz.Clock
|
||||
dialOptions *websocket.DialOptions
|
||||
conn tailnetConn
|
||||
customDialFn func() (proto.DRPCTailnetClient, error)
|
||||
@@ -68,9 +71,10 @@ type tailnetAPIConnector struct {
|
||||
clientMu sync.RWMutex
|
||||
client proto.DRPCTailnetClient
|
||||
|
||||
connected chan error
|
||||
isFirst bool
|
||||
closed chan struct{}
|
||||
connected chan error
|
||||
resumeToken *proto.RefreshResumeTokenResponse
|
||||
isFirst bool
|
||||
closed chan struct{}
|
||||
|
||||
// Only set to true if we get a response from the server that it doesn't support
|
||||
// network telemetry.
|
||||
@@ -78,12 +82,13 @@ type tailnetAPIConnector struct {
|
||||
}
|
||||
|
||||
// Create a new tailnetAPIConnector without running it
|
||||
func newTailnetAPIConnector(ctx context.Context, logger slog.Logger, agentID uuid.UUID, coordinateURL string, dialOptions *websocket.DialOptions) *tailnetAPIConnector {
|
||||
func newTailnetAPIConnector(ctx context.Context, logger slog.Logger, agentID uuid.UUID, coordinateURL string, clock quartz.Clock, dialOptions *websocket.DialOptions) *tailnetAPIConnector {
|
||||
return &tailnetAPIConnector{
|
||||
ctx: ctx,
|
||||
logger: logger,
|
||||
agentID: agentID,
|
||||
coordinateURL: coordinateURL,
|
||||
clock: clock,
|
||||
dialOptions: dialOptions,
|
||||
conn: nil,
|
||||
connected: make(chan error, 1),
|
||||
@@ -96,7 +101,7 @@ func newTailnetAPIConnector(ctx context.Context, logger slog.Logger, agentID uui
|
||||
func (tac *tailnetAPIConnector) manageGracefulTimeout() {
|
||||
defer tac.cancelGracefulCtx()
|
||||
<-tac.ctx.Done()
|
||||
timer := time.NewTimer(tailnetConnectorGracefulTimeout)
|
||||
timer := tac.clock.NewTimer(tailnetConnectorGracefulTimeout, "tailnetAPIClient", "gracefulTimeout")
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-tac.closed:
|
||||
@@ -112,6 +117,8 @@ func (tac *tailnetAPIConnector) runConnector(conn tailnetConn) {
|
||||
go func() {
|
||||
tac.isFirst = true
|
||||
defer close(tac.closed)
|
||||
// Sadly retry doesn't support quartz.Clock yet so this is not
|
||||
// influenced by the configured clock.
|
||||
for retrier := retry.New(50*time.Millisecond, 10*time.Second); retrier.Wait(tac.ctx); {
|
||||
tailnetClient, err := tac.dial()
|
||||
if err != nil {
|
||||
@@ -121,7 +128,7 @@ func (tac *tailnetAPIConnector) runConnector(conn tailnetConn) {
|
||||
tac.client = tailnetClient
|
||||
tac.clientMu.Unlock()
|
||||
tac.logger.Debug(tac.ctx, "obtained tailnet API v2+ client")
|
||||
tac.coordinateAndDERPMap(tailnetClient)
|
||||
tac.runConnectorOnce(tailnetClient)
|
||||
tac.logger.Debug(tac.ctx, "tailnet API v2+ connection lost")
|
||||
}
|
||||
}()
|
||||
@@ -138,8 +145,23 @@ func (tac *tailnetAPIConnector) dial() (proto.DRPCTailnetClient, error) {
|
||||
return tac.customDialFn()
|
||||
}
|
||||
tac.logger.Debug(tac.ctx, "dialing Coder tailnet v2+ API")
|
||||
|
||||
u, err := url.Parse(tac.coordinateURL)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse URL %q: %w", tac.coordinateURL, err)
|
||||
}
|
||||
if tac.resumeToken != nil {
|
||||
q := u.Query()
|
||||
q.Set("resume_token", tac.resumeToken.Token)
|
||||
u.RawQuery = q.Encode()
|
||||
tac.logger.Debug(tac.ctx, "using resume token", slog.F("resume_token", tac.resumeToken))
|
||||
}
|
||||
|
||||
coordinateURL := u.String()
|
||||
tac.logger.Debug(tac.ctx, "using coordinate URL", slog.F("url", coordinateURL))
|
||||
|
||||
// nolint:bodyclose
|
||||
ws, res, err := websocket.Dial(tac.ctx, tac.coordinateURL, tac.dialOptions)
|
||||
ws, res, err := websocket.Dial(tac.ctx, coordinateURL, tac.dialOptions)
|
||||
if tac.isFirst {
|
||||
if res != nil && slices.Contains(permanentErrorStatuses, res.StatusCode) {
|
||||
err = codersdk.ReadBodyAsError(res)
|
||||
@@ -160,8 +182,20 @@ func (tac *tailnetAPIConnector) dial() (proto.DRPCTailnetClient, error) {
|
||||
close(tac.connected)
|
||||
}
|
||||
if err != nil {
|
||||
bodyErr := codersdk.ReadBodyAsError(res)
|
||||
var sdkErr *codersdk.Error
|
||||
if xerrors.As(bodyErr, &sdkErr) {
|
||||
for _, v := range sdkErr.Validations {
|
||||
if v.Field == "resume_token" {
|
||||
// Unset the resume token for the next attempt
|
||||
tac.logger.Warn(tac.ctx, "failed to dial tailnet v2+ API: server replied invalid resume token; unsetting for next connection attempt")
|
||||
tac.resumeToken = nil
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
tac.logger.Error(tac.ctx, "failed to dial tailnet v2+ API", slog.Error(err))
|
||||
tac.logger.Error(tac.ctx, "failed to dial tailnet v2+ API", slog.Error(err), slog.F("sdk_err", sdkErr))
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
@@ -177,11 +211,11 @@ func (tac *tailnetAPIConnector) dial() (proto.DRPCTailnetClient, error) {
|
||||
return client, err
|
||||
}
|
||||
|
||||
// coordinateAndDERPMap uses the provided client to coordinate and stream DERP Maps. It is combined
|
||||
// runConnectorOnce uses the provided client to coordinate and stream DERP Maps. It is combined
|
||||
// into one function so that a problem with one tears down the other and triggers a retry (if
|
||||
// appropriate). We multiplex both RPCs over the same websocket, so we want them to share the same
|
||||
// fate.
|
||||
func (tac *tailnetAPIConnector) coordinateAndDERPMap(client proto.DRPCTailnetClient) {
|
||||
func (tac *tailnetAPIConnector) runConnectorOnce(client proto.DRPCTailnetClient) {
|
||||
defer func() {
|
||||
conn := client.DRPCConn()
|
||||
closeErr := conn.Close()
|
||||
@@ -193,14 +227,17 @@ func (tac *tailnetAPIConnector) coordinateAndDERPMap(client proto.DRPCTailnetCli
|
||||
<-conn.Closed()
|
||||
}
|
||||
}()
|
||||
|
||||
refreshTokenCtx, refreshTokenCancel := context.WithCancel(tac.ctx)
|
||||
wg := sync.WaitGroup{}
|
||||
wg.Add(2)
|
||||
wg.Add(3)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
tac.coordinate(client)
|
||||
}()
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
defer refreshTokenCancel()
|
||||
dErr := tac.derpMap(client)
|
||||
if dErr != nil && tac.ctx.Err() == nil {
|
||||
// The main context is still active, meaning that we want the tailnet data plane to stay
|
||||
@@ -215,6 +252,10 @@ func (tac *tailnetAPIConnector) coordinateAndDERPMap(client proto.DRPCTailnetCli
|
||||
// Note that derpMap() logs it own errors, we don't bother here.
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
tac.refreshToken(refreshTokenCtx, client)
|
||||
}()
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
@@ -278,6 +319,41 @@ func (tac *tailnetAPIConnector) derpMap(client proto.DRPCTailnetClient) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (tac *tailnetAPIConnector) refreshToken(ctx context.Context, client proto.DRPCTailnetClient) {
|
||||
ticker := tac.clock.NewTicker(15*time.Second, "tailnetAPIConnector", "refreshToken")
|
||||
defer ticker.Stop()
|
||||
|
||||
initialCh := make(chan struct{}, 1)
|
||||
initialCh <- struct{}{}
|
||||
defer close(initialCh)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
case <-initialCh:
|
||||
}
|
||||
|
||||
attemptCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
res, err := client.RefreshResumeToken(attemptCtx, &proto.RefreshResumeTokenRequest{})
|
||||
cancel()
|
||||
if err != nil {
|
||||
if ctx.Err() == nil {
|
||||
tac.logger.Error(tac.ctx, "error refreshing coordinator resume token", slog.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
tac.logger.Debug(tac.ctx, "refreshed coordinator resume token", slog.F("resume_token", res))
|
||||
tac.resumeToken = res
|
||||
dur := res.RefreshIn.AsDuration()
|
||||
if dur <= 0 {
|
||||
// A sensible delay to refresh again.
|
||||
dur = 30 * time.Minute
|
||||
}
|
||||
ticker.Reset(dur, "tailnetAPIConnector", "refreshToken", "reset")
|
||||
}
|
||||
}
|
||||
|
||||
func (tac *tailnetAPIConnector) SendTelemetryEvent(event *proto.TelemetryEvent) {
|
||||
tac.clientMu.RLock()
|
||||
// We hold the lock for the entire telemetry request, but this would only block
|
||||
|
||||
@@ -14,6 +14,8 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/xerrors"
|
||||
"google.golang.org/protobuf/types/known/durationpb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
"nhooyr.io/websocket"
|
||||
"storj.io/drpc"
|
||||
"storj.io/drpc/drpcerr"
|
||||
@@ -28,6 +30,7 @@ import (
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/coder/v2/tailnet/tailnettest"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func init() {
|
||||
@@ -59,6 +62,7 @@ func TestTailnetAPIConnector_Disconnects(t *testing.T) {
|
||||
DERPMapUpdateFrequency: time.Millisecond,
|
||||
DERPMapFn: func() *tailcfg.DERPMap { return <-derpMapCh },
|
||||
NetworkTelemetryHandler: func(batch []*proto.TelemetryEvent) {},
|
||||
ResumeTokenProvider: tailnet.NewInsecureTestResumeTokenProvider(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -78,7 +82,7 @@ func TestTailnetAPIConnector_Disconnects(t *testing.T) {
|
||||
|
||||
fConn := newFakeTailnetConn()
|
||||
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, &websocket.DialOptions{})
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, quartz.NewReal(), &websocket.DialOptions{})
|
||||
uut.runConnector(fConn)
|
||||
|
||||
call := testutil.RequireRecvCtx(ctx, t, fCoord.CoordinateCalls)
|
||||
@@ -131,7 +135,7 @@ func TestTailnetAPIConnector_UplevelVersion(t *testing.T) {
|
||||
|
||||
fConn := newFakeTailnetConn()
|
||||
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, &websocket.DialOptions{})
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, quartz.NewReal(), &websocket.DialOptions{})
|
||||
uut.runConnector(fConn)
|
||||
|
||||
err := testutil.RequireRecvCtx(ctx, t, uut.connected)
|
||||
@@ -142,6 +146,215 @@ func TestTailnetAPIConnector_UplevelVersion(t *testing.T) {
|
||||
require.NotEmpty(t, sdkErr.Helper)
|
||||
}
|
||||
|
||||
func TestTailnetAPIConnector_ResumeToken(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := slogtest.Make(t, &slogtest.Options{
|
||||
IgnoreErrors: true,
|
||||
}).Leveled(slog.LevelDebug)
|
||||
agentID := uuid.UUID{0x55}
|
||||
fCoord := tailnettest.NewFakeCoordinator()
|
||||
var coord tailnet.Coordinator = fCoord
|
||||
coordPtr := atomic.Pointer[tailnet.Coordinator]{}
|
||||
coordPtr.Store(&coord)
|
||||
derpMapCh := make(chan *tailcfg.DERPMap)
|
||||
defer close(derpMapCh)
|
||||
|
||||
clock := quartz.NewMock(t)
|
||||
resumeTokenSigningKey, err := tailnet.GenerateResumeTokenSigningKey()
|
||||
require.NoError(t, err)
|
||||
resumeTokenProvider := tailnet.NewResumeTokenKeyProvider(resumeTokenSigningKey, clock, time.Hour)
|
||||
svc, err := tailnet.NewClientService(tailnet.ClientServiceOptions{
|
||||
Logger: logger,
|
||||
CoordPtr: &coordPtr,
|
||||
DERPMapUpdateFrequency: time.Millisecond,
|
||||
DERPMapFn: func() *tailcfg.DERPMap { return <-derpMapCh },
|
||||
NetworkTelemetryHandler: func(batch []*proto.TelemetryEvent) {},
|
||||
ResumeTokenProvider: resumeTokenProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
var (
|
||||
websocketConnCh = make(chan *websocket.Conn, 64)
|
||||
expectResumeToken = ""
|
||||
)
|
||||
svr := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Accept a resume_token query parameter to use the same peer ID. This
|
||||
// behavior matches the actual client coordinate route.
|
||||
var (
|
||||
peerID = uuid.New()
|
||||
resumeToken = r.URL.Query().Get("resume_token")
|
||||
)
|
||||
t.Logf("received resume token: %s", resumeToken)
|
||||
assert.Equal(t, expectResumeToken, resumeToken)
|
||||
if resumeToken != "" {
|
||||
peerID, err = resumeTokenProvider.VerifyResumeToken(resumeToken)
|
||||
assert.NoError(t, err, "failed to parse resume token")
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, w, http.StatusUnauthorized, codersdk.Response{
|
||||
Message: CoordinateAPIInvalidResumeToken,
|
||||
Detail: err.Error(),
|
||||
Validations: []codersdk.ValidationError{
|
||||
{Field: "resume_token", Detail: CoordinateAPIInvalidResumeToken},
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
sws, err := websocket.Accept(w, r, nil)
|
||||
if !assert.NoError(t, err) {
|
||||
return
|
||||
}
|
||||
testutil.RequireSendCtx(ctx, t, websocketConnCh, sws)
|
||||
ctx, nc := codersdk.WebsocketNetConn(r.Context(), sws, websocket.MessageBinary)
|
||||
err = svc.ServeConnV2(ctx, nc, tailnet.StreamID{
|
||||
Name: "client",
|
||||
ID: peerID,
|
||||
Auth: tailnet.ClientCoordinateeAuth{AgentID: agentID},
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
}))
|
||||
|
||||
fConn := newFakeTailnetConn()
|
||||
|
||||
newTickerTrap := clock.Trap().NewTicker("tailnetAPIConnector", "refreshToken")
|
||||
tickerResetTrap := clock.Trap().TickerReset("tailnetAPIConnector", "refreshToken", "reset")
|
||||
defer newTickerTrap.Close()
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, clock, &websocket.DialOptions{})
|
||||
uut.runConnector(fConn)
|
||||
|
||||
// Fetch first token. We don't need to advance the clock since we use a
|
||||
// channel with a single item to immediately fetch.
|
||||
newTickerTrap.MustWait(ctx).Release()
|
||||
// We call ticker.Reset after each token fetch to apply the refresh duration
|
||||
// requested by the server.
|
||||
trappedReset := tickerResetTrap.MustWait(ctx)
|
||||
trappedReset.Release()
|
||||
require.NotNil(t, uut.resumeToken)
|
||||
originalResumeToken := uut.resumeToken.Token
|
||||
|
||||
// Fetch second token.
|
||||
waiter := clock.Advance(trappedReset.Duration)
|
||||
waiter.MustWait(ctx)
|
||||
trappedReset = tickerResetTrap.MustWait(ctx)
|
||||
trappedReset.Release()
|
||||
require.NotNil(t, uut.resumeToken)
|
||||
require.NotEqual(t, originalResumeToken, uut.resumeToken.Token)
|
||||
expectResumeToken = uut.resumeToken.Token
|
||||
t.Logf("expecting resume token: %s", expectResumeToken)
|
||||
|
||||
// Sever the connection and expect it to reconnect with the resume token.
|
||||
wsConn := testutil.RequireRecvCtx(ctx, t, websocketConnCh)
|
||||
_ = wsConn.Close(websocket.StatusGoingAway, "test")
|
||||
|
||||
// Wait for the resume token to be refreshed.
|
||||
trappedTicker := newTickerTrap.MustWait(ctx)
|
||||
// Advance the clock slightly to ensure the new JWT is different.
|
||||
clock.Advance(time.Second).MustWait(ctx)
|
||||
trappedTicker.Release()
|
||||
trappedReset = tickerResetTrap.MustWait(ctx)
|
||||
trappedReset.Release()
|
||||
|
||||
// The resume token should have changed again.
|
||||
require.NotNil(t, uut.resumeToken)
|
||||
require.NotEqual(t, expectResumeToken, uut.resumeToken.Token)
|
||||
}
|
||||
|
||||
func TestTailnetAPIConnector_ResumeTokenFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
logger := slogtest.Make(t, &slogtest.Options{
|
||||
IgnoreErrors: true,
|
||||
}).Leveled(slog.LevelDebug)
|
||||
agentID := uuid.UUID{0x55}
|
||||
fCoord := tailnettest.NewFakeCoordinator()
|
||||
var coord tailnet.Coordinator = fCoord
|
||||
coordPtr := atomic.Pointer[tailnet.Coordinator]{}
|
||||
coordPtr.Store(&coord)
|
||||
derpMapCh := make(chan *tailcfg.DERPMap)
|
||||
defer close(derpMapCh)
|
||||
|
||||
clock := quartz.NewMock(t)
|
||||
resumeTokenSigningKey, err := tailnet.GenerateResumeTokenSigningKey()
|
||||
require.NoError(t, err)
|
||||
resumeTokenProvider := tailnet.NewResumeTokenKeyProvider(resumeTokenSigningKey, clock, time.Hour)
|
||||
svc, err := tailnet.NewClientService(tailnet.ClientServiceOptions{
|
||||
Logger: logger,
|
||||
CoordPtr: &coordPtr,
|
||||
DERPMapUpdateFrequency: time.Millisecond,
|
||||
DERPMapFn: func() *tailcfg.DERPMap { return <-derpMapCh },
|
||||
NetworkTelemetryHandler: func(batch []*proto.TelemetryEvent) {},
|
||||
ResumeTokenProvider: resumeTokenProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
var (
|
||||
websocketConnCh = make(chan *websocket.Conn, 64)
|
||||
didFail int64
|
||||
)
|
||||
svr := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Query().Get("resume_token") != "" {
|
||||
atomic.AddInt64(&didFail, 1)
|
||||
httpapi.Write(ctx, w, http.StatusUnauthorized, codersdk.Response{
|
||||
Message: CoordinateAPIInvalidResumeToken,
|
||||
Validations: []codersdk.ValidationError{
|
||||
{Field: "resume_token", Detail: CoordinateAPIInvalidResumeToken},
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
sws, err := websocket.Accept(w, r, nil)
|
||||
if !assert.NoError(t, err) {
|
||||
return
|
||||
}
|
||||
testutil.RequireSendCtx(ctx, t, websocketConnCh, sws)
|
||||
ctx, nc := codersdk.WebsocketNetConn(r.Context(), sws, websocket.MessageBinary)
|
||||
err = svc.ServeConnV2(ctx, nc, tailnet.StreamID{
|
||||
Name: "client",
|
||||
ID: uuid.New(),
|
||||
Auth: tailnet.ClientCoordinateeAuth{AgentID: agentID},
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
}))
|
||||
|
||||
fConn := newFakeTailnetConn()
|
||||
|
||||
newTickerTrap := clock.Trap().NewTicker("tailnetAPIConnector", "refreshToken")
|
||||
tickerResetTrap := clock.Trap().TickerReset("tailnetAPIConnector", "refreshToken", "reset")
|
||||
defer newTickerTrap.Close()
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, clock, &websocket.DialOptions{})
|
||||
uut.runConnector(fConn)
|
||||
|
||||
// Wait for the resume token to be fetched for the first time.
|
||||
newTickerTrap.MustWait(ctx).Release()
|
||||
trappedReset := tickerResetTrap.MustWait(ctx)
|
||||
trappedReset.Release()
|
||||
originalResumeToken := uut.resumeToken.Token
|
||||
|
||||
// Sever the connection and expect it to reconnect with the resume token,
|
||||
// which should fail and cause the client to be disconnected. The client
|
||||
// should then reconnect with no resume token.
|
||||
wsConn := testutil.RequireRecvCtx(ctx, t, websocketConnCh)
|
||||
_ = wsConn.Close(websocket.StatusGoingAway, "test")
|
||||
|
||||
// Wait for the resume token to be refreshed, which indicates a successful
|
||||
// reconnect.
|
||||
trappedTicker := newTickerTrap.MustWait(ctx)
|
||||
// Since we failed the initial reconnect and we're definitely reconnected
|
||||
// now, the stored resume token should now be nil.
|
||||
require.Nil(t, uut.resumeToken)
|
||||
trappedTicker.Release()
|
||||
trappedReset = tickerResetTrap.MustWait(ctx)
|
||||
trappedReset.Release()
|
||||
require.NotNil(t, uut.resumeToken)
|
||||
require.NotEqual(t, originalResumeToken, uut.resumeToken.Token)
|
||||
|
||||
// The resume token should have been rejected by the server.
|
||||
require.EqualValues(t, 1, atomic.LoadInt64(&didFail))
|
||||
}
|
||||
|
||||
func TestTailnetAPIConnector_TelemetrySuccess(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
@@ -161,8 +374,9 @@ func TestTailnetAPIConnector_TelemetrySuccess(t *testing.T) {
|
||||
DERPMapUpdateFrequency: time.Millisecond,
|
||||
DERPMapFn: func() *tailcfg.DERPMap { return <-derpMapCh },
|
||||
NetworkTelemetryHandler: func(batch []*proto.TelemetryEvent) {
|
||||
eventCh <- batch
|
||||
testutil.RequireSendCtx(ctx, t, eventCh, batch)
|
||||
},
|
||||
ResumeTokenProvider: tailnet.NewInsecureTestResumeTokenProvider(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -182,7 +396,7 @@ func TestTailnetAPIConnector_TelemetrySuccess(t *testing.T) {
|
||||
|
||||
fConn := newFakeTailnetConn()
|
||||
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, &websocket.DialOptions{})
|
||||
uut := newTailnetAPIConnector(ctx, logger, agentID, svr.URL, quartz.NewReal(), &websocket.DialOptions{})
|
||||
uut.runConnector(fConn)
|
||||
require.Eventually(t, func() bool {
|
||||
uut.clientMu.Lock()
|
||||
@@ -213,6 +427,7 @@ func TestTailnetAPIConnector_TelemetryUnimplemented(t *testing.T) {
|
||||
logger: logger,
|
||||
agentID: agentID,
|
||||
coordinateURL: "",
|
||||
clock: quartz.NewReal(),
|
||||
dialOptions: &websocket.DialOptions{},
|
||||
conn: nil,
|
||||
connected: make(chan error, 1),
|
||||
@@ -253,6 +468,7 @@ func TestTailnetAPIConnector_TelemetryNotRecognised(t *testing.T) {
|
||||
logger: logger,
|
||||
agentID: agentID,
|
||||
coordinateURL: "",
|
||||
clock: quartz.NewReal(),
|
||||
dialOptions: &websocket.DialOptions{},
|
||||
conn: nil,
|
||||
connected: make(chan error, 1),
|
||||
@@ -301,6 +517,7 @@ func newFakeTailnetConn() *fakeTailnetConn {
|
||||
|
||||
type fakeDRPCClient struct {
|
||||
postTelemetryCalls int64
|
||||
refreshTokenFn func(context.Context, *proto.RefreshResumeTokenRequest) (*proto.RefreshResumeTokenResponse, error)
|
||||
telemetryError error
|
||||
fakeDRPPCMapStream
|
||||
}
|
||||
@@ -339,6 +556,19 @@ func (f *fakeDRPCClient) StreamDERPMaps(_ context.Context, _ *proto.StreamDERPMa
|
||||
return &f.fakeDRPPCMapStream, nil
|
||||
}
|
||||
|
||||
// RefreshResumeToken implements proto.DRPCTailnetClient.
|
||||
func (f *fakeDRPCClient) RefreshResumeToken(_ context.Context, _ *proto.RefreshResumeTokenRequest) (*proto.RefreshResumeTokenResponse, error) {
|
||||
if f.refreshTokenFn != nil {
|
||||
return f.refreshTokenFn(context.Background(), nil)
|
||||
}
|
||||
|
||||
return &proto.RefreshResumeTokenResponse{
|
||||
Token: "test",
|
||||
RefreshIn: durationpb.New(30 * time.Minute),
|
||||
ExpiresAt: timestamppb.New(time.Now().Add(time.Hour)),
|
||||
}, nil
|
||||
}
|
||||
|
||||
type fakeDRPCConn struct{}
|
||||
|
||||
var _ drpc.Conn = &fakeDRPCConn{}
|
||||
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
// AgentIP is a static IPv6 address with the Tailscale prefix that is used to route
|
||||
@@ -55,7 +56,11 @@ const (
|
||||
AgentMinimumListeningPort = 9
|
||||
)
|
||||
|
||||
const AgentAPIMismatchMessage = "Unknown or unsupported API version"
|
||||
const (
|
||||
AgentAPIMismatchMessage = "Unknown or unsupported API version"
|
||||
|
||||
CoordinateAPIInvalidResumeToken = "Invalid resume token"
|
||||
)
|
||||
|
||||
// AgentIgnoredListeningPorts contains a list of ports to ignore when looking for
|
||||
// running applications inside a workspace. We want to ignore non-HTTP servers,
|
||||
@@ -232,7 +237,7 @@ func (c *Client) DialAgent(dialCtx context.Context, agentID uuid.UUID, options *
|
||||
q.Add("version", "2.0")
|
||||
coordinateURL.RawQuery = q.Encode()
|
||||
|
||||
connector := newTailnetAPIConnector(ctx, options.Logger, agentID, coordinateURL.String(),
|
||||
connector := newTailnetAPIConnector(ctx, options.Logger, agentID, coordinateURL.String(), quartz.NewReal(),
|
||||
&websocket.DialOptions{
|
||||
HTTPClient: c.client.HTTPClient,
|
||||
HTTPHeader: headers,
|
||||
|
||||
Reference in New Issue
Block a user