diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index 7364cd269b6..3bec604822b 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -421,7 +421,8 @@ func bedrockModelSupportsAdaptiveThinking(model string) bool { // // See https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-opus-4-7.html func bedrockModelRequiresAdaptiveThinking(model string) bool { - return strings.Contains(model, "anthropic.claude-opus-4-7") + return strings.Contains(model, "anthropic.claude-opus-4-7") || + strings.Contains(model, "anthropic.claude-opus-4-8") } // filterBedrockBetaFlags removes unsupported beta flags from the Anthropic-Beta diff --git a/aibridge/intercept/messages/base_test.go b/aibridge/intercept/messages/base_test.go index e4314486546..735ea9c6d27 100644 --- a/aibridge/intercept/messages/base_test.go +++ b/aibridge/intercept/messages/base_test.go @@ -846,6 +846,12 @@ func TestAugmentRequestForBedrock_AdaptiveThinking(t *testing.T) { requestBody: `{"max_tokens":10000,"thinking":{"type":"enabled","budget_tokens":8000}}`, expectThinkingType: "adaptive", }, + { + name: "opus_4_8_model_with_enabled_thinking_is_converted_to_adaptive_and_drops_budget", + bedrockModel: "eu.anthropic.claude-opus-4-8", + requestBody: `{"max_tokens":10000,"thinking":{"type":"enabled","budget_tokens":5000}}`, + expectThinkingType: "adaptive", + }, { // Opus 4.7 on Bedrock rejects output_config.format (structured // outputs) with a 400 even though it accepts output_config.effort. diff --git a/coderd/aitasks.go b/coderd/aitasks.go index f0adb3b8ea8..5b45828be4f 100644 --- a/coderd/aitasks.go +++ b/coderd/aitasks.go @@ -6,7 +6,6 @@ import ( "encoding/json" "errors" "fmt" - "net" "net/http" "net/url" "slices" @@ -1086,13 +1085,7 @@ func (api *API) authAndDoWithTaskAppClient( } defer release() - client := &http.Client{ - Transport: &http.Transport{ - DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { - return agentConn.DialContext(ctx, network, addr) - }, - }, - } + client := agentConn.AppHTTPClient() return do(ctx, client, parsedURL) } diff --git a/coderd/azureidentity/azureidentity_test.go b/coderd/azureidentity/azureidentity_test.go index 9ed2750f454..c10bbfbe803 100644 --- a/coderd/azureidentity/azureidentity_test.go +++ b/coderd/azureidentity/azureidentity_test.go @@ -21,6 +21,7 @@ import ( func TestValidate(t *testing.T) { t.Parallel() + t.Skip("See https://github.com/coder/internal/issues/1602") if runtime.GOOS == "darwin" { // This test fails on MacOS for some reason. See https://github.com/coder/coder/issues/12978 t.Skip() diff --git a/coderd/database/dbfake/dbfake.go b/coderd/database/dbfake/dbfake.go index 0b859a4fb1c..82b66f504aa 100644 --- a/coderd/database/dbfake/dbfake.go +++ b/coderd/database/dbfake/dbfake.go @@ -69,6 +69,8 @@ type WorkspaceBuildBuilder struct { jobErrorCode string // Error code for failed jobs provisionerState []byte + + prebuiltWorkspaceBuildStage sdkproto.PrebuiltWorkspaceBuildStage } // BuilderOption is a functional option for customizing job timestamps @@ -149,6 +151,14 @@ func (b WorkspaceBuildBuilder) ProvisionerState(state []byte) WorkspaceBuildBuil return b } +// MarkPrebuiltWorkspaceClaim marks the build's provisioner job as the claim +// of a prebuilt workspace, mirroring wsbuilder.MarkPrebuiltWorkspaceClaim. +func (b WorkspaceBuildBuilder) MarkPrebuiltWorkspaceClaim() WorkspaceBuildBuilder { + //nolint: revive // returns modified struct + b.prebuiltWorkspaceBuildStage = sdkproto.PrebuiltWorkspaceBuildStage_CLAIM + return b +} + func (b WorkspaceBuildBuilder) Resource(resource ...*sdkproto.Resource) WorkspaceBuildBuilder { //nolint: revive // returns modified struct b.resources = append(b.resources, resource...) @@ -368,7 +378,8 @@ func (b WorkspaceBuildBuilder) doInTX() WorkspaceResponse { // Create a provisioner job for the build! payload, err := json.Marshal(provisionerdserver.WorkspaceProvisionJob{ - WorkspaceBuildID: b.seed.ID, + WorkspaceBuildID: b.seed.ID, + PrebuiltWorkspaceBuildStage: b.prebuiltWorkspaceBuildStage, }) require.NoError(b.t, err) diff --git a/coderd/tailnet.go b/coderd/tailnet.go index 6f591835d94..0ae17cfaef7 100644 --- a/coderd/tailnet.go +++ b/coderd/tailnet.go @@ -298,6 +298,7 @@ func (s *ServerTailnet) AgentConn(ctx context.Context, agentID uuid.UUID) (works conn = workspacesdk.NewAgentConn(s.conn, workspacesdk.AgentConnOptions{ AgentID: agentID, CloseFunc: func() error { return workspacesdk.ErrSkipClose }, + Logger: s.logger, }) // Since we now have an open conn, be careful to close it if we error diff --git a/coderd/workspaceagents.go b/coderd/workspaceagents.go index bfd7a8e7867..89f574d9f24 100644 --- a/coderd/workspaceagents.go +++ b/coderd/workspaceagents.go @@ -37,6 +37,7 @@ import ( "github.com/coder/coder/v2/coderd/httpmw/loggermw" "github.com/coder/coder/v2/coderd/jwtutils" "github.com/coder/coder/v2/coderd/prebuilds" + "github.com/coder/coder/v2/coderd/provisionerdserver" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/rbac/policy" "github.com/coder/coder/v2/coderd/telemetry" @@ -1541,14 +1542,13 @@ func (api *API) workspaceAgentReinit(rw http.ResponseWriter, r *http.Request) { return } - // This workspace was a prebuild that got claimed. Check if - // the claim build completed successfully before sending - // reinit. We assume the latest build is the claim build - // (build 2). If a third build (e.g. a restart) starts - // between the claim and the agent's reconnection, this - // would check that build instead. The window is extremely - // small in practice, and a restart would trigger its own - // reinit path. + // This workspace was a prebuild that got claimed. The seeded + // reinit below recovers a claim event that was missed while + // the agent's /reinit connection was down. It only applies + // while the latest build is the claim build itself, which the + // build's provisioner job input records, mirroring the check + // the provisioner server uses when publishing the claim + // event. latestBuild, err := api.Database.GetLatestWorkspaceBuildByWorkspaceID(ctx, workspace.ID) if err != nil { log.Error(ctx, "failed to get latest workspace build", slog.Error(err)) @@ -1561,43 +1561,73 @@ func (api *API) workspaceAgentReinit(rw http.ResponseWriter, r *http.Request) { httpapi.InternalServerError(rw, xerrors.New("failed to get provisioner job")) return } + var jobInput provisionerdserver.WorkspaceProvisionJob + if err := json.Unmarshal(job.Input, &jobInput); err != nil { + log.Error(ctx, "failed to unmarshal provisioner job input", slog.Error(err)) + httpapi.InternalServerError(rw, xerrors.New("failed to unmarshal provisioner job input")) + return + } - if job.CompletedAt.Valid && !job.Error.Valid { - // Claim build succeeded — cancel the pubsub - // subscription (no longer needed) and swap in a - // pre-seeded channel so the transmitter delivers - // exactly one reinit event. - cancelSub() - seeded := make(chan agentsdk.ReinitializationEvent, 1) - seeded <- agentsdk.ReinitializationEvent{ - WorkspaceID: workspace.ID, - Reason: agentsdk.ReinitializeReasonPrebuildClaimed, - OwnerID: workspace.OwnerID, + switch { + case jobInput.PrebuiltWorkspaceBuildStage.IsPrebuiltWorkspaceClaim(): + if job.CompletedAt.Valid && !job.Error.Valid { + // Claim build succeeded: cancel the pubsub + // subscription (no longer needed) and swap in a + // pre-seeded channel so the transmitter delivers + // exactly one reinit event. + cancelSub() + seeded := make(chan agentsdk.ReinitializationEvent, 1) + seeded <- agentsdk.ReinitializationEvent{ + WorkspaceID: workspace.ID, + Reason: agentsdk.ReinitializeReasonPrebuildClaimed, + OwnerID: workspace.OwnerID, + } + reinitEvents = seeded + } else if job.CompletedAt.Valid && job.Error.Valid { + // Claim build failed permanently. Return 409 so the + // agent treats this as terminal and stops retrying + // (WaitForReinitLoop exits on any 409). + cancelSub() + log.Warn(ctx, "claim build failed", + slog.F("job_id", job.ID), + slog.F("error", job.Error.String)) + httpapi.Write(ctx, rw, http.StatusConflict, codersdk.Response{ + Message: "Claim build failed permanently.", + Detail: job.Error.String, + }) + return } - reinitEvents = seeded - } else if job.CompletedAt.Valid && job.Error.Valid { - // Claim build failed permanently. Return 409 so the - // agent treats this as terminal and stops retrying - // (WaitForReinitLoop exits on any 409). - cancelSub() - log.Warn(ctx, "claim build failed", + // Claim build still in progress: proceed to the + // transmitter below. The pubsub subscription (set up + // above) will deliver the event when the build completes + // successfully. Note: FailJob does not publish a claim + // event, so a failed in-progress build will leave the + // agent blocking here until it disconnects and + // reconnects (at which point the durable check above + // handles it). + case latestBuild.InitiatorID == database.PrebuildsSystemUserID: + // The workspace owner has changed but the claim build has + // not been created yet. Proceed to the transmitter below; + // the pubsub subscription set up above delivers the claim + // event once the claim build completes. + default: + // The latest build is a user-initiated build other than + // the claim build, so the claim has already been handled. + // Re-sending the reinit event here would needlessly + // restart the agent of a long-claimed workspace on every + // reconnection. Return 409 so the agent stops polling, + // the same as a regular workspace. + log.Debug(ctx, "prebuild claim already handled, stopping reinit polling", slog.F("job_id", job.ID), - slog.F("error", job.Error.String)) + slog.F("latest_build_id", latestBuild.ID), + slog.F("latest_build_number", latestBuild.BuildNumber)) + cancelSub() httpapi.Write(ctx, rw, http.StatusConflict, codersdk.Response{ - Message: "Claim build failed permanently.", - Detail: job.Error.String, + Message: "Workspace is not a prebuilt workspace waiting to be claimed.", + Detail: "The prebuild claim for this workspace has already been handled by an earlier build.", }) return } - - // Claim build still in progress — fall through to the - // transmitter. The pubsub subscription (set up above) - // will deliver the event when the build completes - // successfully. Note: FailJob does not publish a claim - // event, so a failed in-progress build will leave the - // agent blocking here until it disconnects and - // reconnects (at which point the durable check above - // handles it). } transmitter := agentsdk.NewSSEAgentReinitTransmitter(log, rw, r) diff --git a/coderd/workspaceagents_test.go b/coderd/workspaceagents_test.go index cdb265686c8..add08d8aa4a 100644 --- a/coderd/workspaceagents_test.go +++ b/coderd/workspaceagents_test.go @@ -3399,6 +3399,7 @@ func TestReinit(t *testing.T) { InitiatorID: claimerID, Transition: database.WorkspaceTransitionStart, }). + MarkPrebuiltWorkspaceClaim(). WithAgent() if !complete { builder = builder.Starting() @@ -3496,6 +3497,52 @@ func TestReinit(t *testing.T) { require.Equal(t, user.UserID, reinitEvent.OwnerID) }) + // Verifies that the durable claim check only applies while the + // latest build is the claim build. A workspace that was claimed + // in the past and has since had user-initiated builds must get a + // 409 instead of another reinit, otherwise its agent would be + // restarted on every /reinit reconnection for the rest of the + // workspace's life. + t.Run("workspace claimed in the past gets 409", func(t *testing.T) { + t.Parallel() + + db, ps, sqlDB := dbtestutil.NewDBWithSQLDB(t) + client := coderdtest.New(t, &coderdtest.Options{ + Database: db, + Pubsub: ps, + }) + user := coderdtest.CreateFirstUser(t, client) + + // Create an unclaimed prebuild (build 1, completed) and claim + // it (build 2, completed). + r := setupPrebuildWorkspace(t, db, user.OrganizationID) + claimPrebuild(t, db, sqlDB, r.Workspace, user.UserID, r.TemplateVersion.ID, true) + + // A later build initiated by the owner (e.g. a restart) means + // the claim has already been handled. + ws := r.Workspace + ws.OwnerID = user.UserID + laterR := dbfake.WorkspaceBuild(t, db, ws). + Seed(database.WorkspaceBuild{ + TemplateVersionID: r.TemplateVersion.ID, + BuildNumber: 3, + InitiatorID: user.UserID, + Transition: database.WorkspaceTransitionStart, + }). + WithAgent(). + Do() + + agentCtx := testutil.Context(t, testutil.WaitShort) + agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(laterR.AgentToken)) + + // WaitForReinit should return an error wrapping a 409. + _, err := agentClient.WaitForReinit(agentCtx) + require.Error(t, err) + var sdkErr *codersdk.Error + require.ErrorAs(t, err, &sdkErr) + require.Equal(t, http.StatusConflict, sdkErr.StatusCode()) + }) + // Verifies that when the claim build completed with an error, // the handler returns 409 so the agent treats it as terminal // and stops retrying (WaitForReinitLoop exits on any 409). diff --git a/codersdk/workspacesdk/agentconn.go b/codersdk/workspacesdk/agentconn.go index 6882ff0d916..2f005c9c040 100644 --- a/codersdk/workspacesdk/agentconn.go +++ b/codersdk/workspacesdk/agentconn.go @@ -102,6 +102,7 @@ type AgentConn interface { DebugMagicsock(ctx context.Context) ([]byte, error) DebugManifest(ctx context.Context) ([]byte, error) DialContext(ctx context.Context, network string, addr string) (net.Conn, error) + AppHTTPClient() *http.Client GetPeerDiagnostics() tailnet.PeerDiagnostics ListContainers(ctx context.Context) (codersdk.WorkspaceAgentListContainersResponse, error) ListMCPTools(ctx context.Context) (ListMCPToolsResponse, error) @@ -158,6 +159,7 @@ func (c *agentConn) SetExtraHeaders(h http.Header) { type AgentConnOptions struct { AgentID uuid.UUID CloseFunc func() error + Logger slog.Logger } func (c *agentConn) agentAddress() netip.Addr { @@ -370,6 +372,24 @@ func (c *agentConn) DialContext(ctx context.Context, network string, addr string } } +// AppHTTPClient returns an HTTP client for reaching HTTP apps served by this +// workspace agent. Redirects are blocked to prevent misuse. +func (c *agentConn) AppHTTPClient() *http.Client { + return &http.Client{ + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + Transport: &http.Transport{ + // Disable keep-alives so these short-lived clients don't leave + // idle connections (and their goroutines) lingering after they're + // discarded. + DisableKeepAlives: true, + // Host locked to agent, port from URL. + DialContext: c.DialContext, + }, + } +} + // ListeningPorts lists the ports that are currently in use by the workspace. func (c *agentConn) ListeningPorts(ctx context.Context) (codersdk.WorkspaceAgentListeningPortsResponse, error) { ctx, span := tracing.StartSpan(ctx) @@ -1358,7 +1378,12 @@ func (c *agentConn) apiRequest(ctx context.Context, method, path string, body in // apiClient returns an HTTP client that can be used to make // requests to the workspace agent's HTTP API server. func (c *agentConn) apiClient() *http.Client { + agentAddr := netip.AddrPortFrom(c.agentAddress(), AgentHTTPAPIServerPort) return &http.Client{ + // Redirects are blocked to prevent misuse. + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, Transport: &http.Transport{ // Disable keep alives as we're usually only making a single // request, and this triggers goleak in tests @@ -1378,16 +1403,21 @@ func (c *agentConn) apiClient() *http.Client { return nil, xerrors.Errorf("request %q does not appear to be for http api", addr) } - if !c.AwaitReachable(ctx) { - return nil, xerrors.Errorf("workspace agent not reachable in time: %v", ctx.Err()) + if reqAddr, err := netip.ParseAddr(host); err != nil || reqAddr != agentAddr.Addr() { + c.opts.Logger.Warn(ctx, "blocked workspace agent API request to unintended host", + slog.F("agent_id", c.opts.AgentID), + slog.F("request_host", host), + slog.F("intended_agent_addr", agentAddr.Addr()), + ) + return nil, xerrors.Errorf("request host %q does not match intended agent %q", host, agentAddr.Addr()) } - ipAddr, err := netip.ParseAddr(host) - if err != nil { - return nil, xerrors.Errorf("parse host addr: %w", err) + if !c.AwaitReachable(ctx) { + return nil, xerrors.Errorf("workspace agent not reachable in time: %v", ctx.Err()) } - conn, err := c.Conn.DialContextTCP(ctx, netip.AddrPortFrom(ipAddr, AgentHTTPAPIServerPort)) + // Always dial the pinned agent address, never the request host. + conn, err := c.Conn.DialContextTCP(ctx, agentAddr) if err != nil { return nil, xerrors.Errorf("dial http api: %w", err) } diff --git a/codersdk/workspacesdk/agentconn_redirect_test.go b/codersdk/workspacesdk/agentconn_redirect_test.go new file mode 100644 index 00000000000..7407be46f22 --- /dev/null +++ b/codersdk/workspacesdk/agentconn_redirect_test.go @@ -0,0 +1,212 @@ +package workspacesdk_test + +import ( + "context" + "errors" + "fmt" + "net" + "net/http" + "net/netip" + "strings" + "sync/atomic" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "tailscale.com/tailcfg" + + "github.com/coder/coder/v2/codersdk/workspacesdk" + "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" +) + +func TestAgentConnRejectsCrossAgentRedirects(t *testing.T) { + t.Parallel() + + derpMap, _ := tailnettest.RunDERPAndSTUN(t) + cases := []struct { + name string + status int + invoke func(context.Context, workspacesdk.AgentConn) error + }{ + { + name: "get 302", + status: http.StatusFound, + invoke: func(ctx context.Context, conn workspacesdk.AgentConn) error { + _, err := conn.ListeningPorts(ctx) + return err + }, + }, + { + name: "post 307", + status: http.StatusTemporaryRedirect, + invoke: func(ctx context.Context, conn workspacesdk.AgentConn) error { + return conn.WriteFile(ctx, "/tmp/attacker", strings.NewReader("redirect-body")) + }, + }, + { + name: "post 308", + status: http.StatusPermanentRedirect, + invoke: func(ctx context.Context, conn workspacesdk.AgentConn) error { + return conn.WriteFile(ctx, "/tmp/attacker", strings.NewReader("redirect-body")) + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitMedium) + + clientID := uuid.New() + attackerID := uuid.New() + victimID := uuid.New() + clientConn, _ := newTailnetConn(t, derpMap, clientID, "client") + attackerConn, attackerIP := newTailnetConn(t, derpMap, attackerID, "attacker") + victimConn, victimIP := newTailnetConn(t, derpMap, victimID, "victim") + stitchTailnet(t, map[uuid.UUID]*tailnet.Conn{ + clientID: clientConn, + attackerID: attackerConn, + victimID: victimConn, + }) + + var victimHit atomic.Bool + victimRouter := http.NewServeMux() + victimRouter.HandleFunc("/api/v0/listening-ports", func(rw http.ResponseWriter, _ *http.Request) { + victimHit.Store(true) + rw.Header().Set("Content-Type", "application/json") + _, _ = rw.Write([]byte(`{"ports":[]}`)) + }) + victimRouter.HandleFunc("/api/v0/write-file", func(rw http.ResponseWriter, _ *http.Request) { + victimHit.Store(true) + rw.WriteHeader(http.StatusOK) + }) + serveTailnetHTTP(t, victimConn, victimRouter) + + victimBaseURL := fmt.Sprintf("http://[%s]:%d", victimIP, workspacesdk.AgentHTTPAPIServerPort) + attackerRouter := http.NewServeMux() + attackerRouter.HandleFunc("/", func(rw http.ResponseWriter, r *http.Request) { + http.Redirect(rw, r, victimBaseURL+r.URL.RequestURI(), tc.status) + }) + serveTailnetHTTP(t, attackerConn, attackerRouter) + + require.True(t, clientConn.AwaitReachable(ctx, attackerIP)) + require.True(t, clientConn.AwaitReachable(ctx, victimIP)) + + conn := workspacesdk.NewAgentConn(clientConn, workspacesdk.AgentConnOptions{ + AgentID: attackerID, + }) + + err := tc.invoke(ctx, conn) + require.Error(t, err) + require.False(t, victimHit.Load()) + }) + } +} + +// TestAgentConnAppHTTPClientRefusesRedirects verifies the app HTTP client does +// not follow redirects. +func TestAgentConnAppHTTPClientRefusesRedirects(t *testing.T) { + t.Parallel() + + tailnetConn, err := tailnet.NewConn(&tailnet.Options{ + Addresses: []netip.Prefix{tailnet.TailscaleServicePrefix.RandomPrefix()}, + Logger: testutil.Logger(t), + }) + require.NoError(t, err) + t.Cleanup(func() { + _ = tailnetConn.Close() + }) + + conn := workspacesdk.NewAgentConn(tailnetConn, workspacesdk.AgentConnOptions{ + AgentID: uuid.New(), + }) + + client := conn.AppHTTPClient() + require.NotNil(t, client.CheckRedirect) + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://example.invalid/", nil) + require.NoError(t, err) + require.ErrorIs(t, client.CheckRedirect(req, nil), http.ErrUseLastResponse) +} + +func newTailnetConn(t *testing.T, derpMap *tailcfg.DERPMap, id uuid.UUID, name string) (*tailnet.Conn, netip.Addr) { + t.Helper() + + addr := tailnet.TailscaleServicePrefix.AddrFromUUID(id) + conn, err := tailnet.NewConn(&tailnet.Options{ + ID: id, + Addresses: []netip.Prefix{netip.PrefixFrom(addr, 128)}, + Logger: testutil.Logger(t).Named(name), + DERPMap: derpMap, + }) + require.NoError(t, err) + t.Cleanup(func() { + assert.NoError(t, conn.Close()) + }) + + return conn, addr +} + +func serveTailnetHTTP(t *testing.T, conn *tailnet.Conn, handler http.Handler) { + t.Helper() + + ln, err := conn.Listen("tcp", fmt.Sprintf(":%d", workspacesdk.AgentHTTPAPIServerPort)) + require.NoError(t, err) + + server := &http.Server{Handler: handler, ReadHeaderTimeout: testutil.WaitShort} + t.Cleanup(func() { + assert.NoError(t, server.Close()) + assert.NoError(t, ln.Close()) + }) + + go func() { + err := server.Serve(ln) + if err != nil && !errors.Is(err, net.ErrClosed) && !errors.Is(err, http.ErrServerClosed) { + assert.NoError(t, err) + } + }() +} + +// stitchTailnet cross-programs every conn's node into every other conn, the +// N-peer analog of tailnet's stitch test helper, so the peers can reach each +// other without a coordinator. +func stitchTailnet(t *testing.T, conns map[uuid.UUID]*tailnet.Conn) { + t.Helper() + + sendNode := func(srcID uuid.UUID, node *tailnet.Node) { + protoNode, err := tailnet.NodeToProto(node) + if !assert.NoError(t, err) { + return + } + for dstID, dst := range conns { + if dstID == srcID { + continue + } + err = dst.UpdatePeers([]*proto.CoordinateResponse_PeerUpdate{{ + Id: srcID[:], + Node: protoNode, + Kind: proto.CoordinateResponse_PeerUpdate_NODE, + }}) + assert.NoError(t, err) + } + } + + for srcID, src := range conns { + src.SetNodeCallback(func(node *tailnet.Node) { + sendNode(srcID, node) + }) + if node := src.Node(); node != nil { + sendNode(srcID, node) + } + } + + t.Cleanup(func() { + for _, conn := range conns { + conn.SetNodeCallback(nil) + } + }) +} diff --git a/codersdk/workspacesdk/agentconnmock/agentconnmock.go b/codersdk/workspacesdk/agentconnmock/agentconnmock.go index 5c23246cae8..d74bee773c6 100644 --- a/codersdk/workspacesdk/agentconnmock/agentconnmock.go +++ b/codersdk/workspacesdk/agentconnmock/agentconnmock.go @@ -56,6 +56,20 @@ func (m *MockAgentConn) EXPECT() *MockAgentConnMockRecorder { return m.recorder } +// AppHTTPClient mocks base method. +func (m *MockAgentConn) AppHTTPClient() *http.Client { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "AppHTTPClient") + ret0, _ := ret[0].(*http.Client) + return ret0 +} + +// AppHTTPClient indicates an expected call of AppHTTPClient. +func (mr *MockAgentConnMockRecorder) AppHTTPClient() *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AppHTTPClient", reflect.TypeOf((*MockAgentConn)(nil).AppHTTPClient)) +} + // AwaitReachable mocks base method. func (m *MockAgentConn) AwaitReachable(ctx context.Context) bool { m.ctrl.T.Helper() diff --git a/codersdk/workspacesdk/workspacesdk.go b/codersdk/workspacesdk/workspacesdk.go index 67eab8b4bcb..f0311c7973f 100644 --- a/codersdk/workspacesdk/workspacesdk.go +++ b/codersdk/workspacesdk/workspacesdk.go @@ -300,6 +300,7 @@ func (c *Client) DialAgent(dialCtx context.Context, agentID uuid.UUID, options * <-controller.Closed() return conn.Close() }, + Logger: options.Logger, }) if !agentConn.AwaitReachable(dialCtx) { diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd.go b/enterprise/aibridgeproxyd/aibridgeproxyd.go index f5048b6fd30..10037728962 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd.go @@ -112,13 +112,15 @@ var blockedIPRanges = func() []net.IPNet { // - decrypting requests using the configured MITM CA certificate // - forwarding requests to aibridged for processing type Server struct { - ctx context.Context - logger slog.Logger - proxy *goproxy.ProxyHttpServer - httpServer *http.Server - listener net.Listener - tlsEnabled bool - coderAccessURL *url.URL + ctx context.Context + logger slog.Logger + proxy *goproxy.ProxyHttpServer + httpServer *http.Server + listener net.Listener + tlsEnabled bool + coderAccessURL *url.URL + // coderAccessPort is the resolved port for the Coder access URL. + coderAccessPort string aibridgeProviderFromHost func(host string) string // caCert is the PEM-encoded MITM CA certificate loaded during initialization. // This is served to clients who need to trust the proxy's generated certificates. @@ -228,7 +230,6 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) coderAccessPort = "80" } } - coderAccessURL.Host = net.JoinHostPort(coderAccessURL.Hostname(), coderAccessPort) // MITM cert and key are required to intercept and decrypt HTTPS traffic. if opts.MITMCertFile == "" || opts.MITMKeyFile == "" { @@ -311,6 +312,7 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) proxy: proxy, tlsEnabled: opts.TLSCertFile != "", coderAccessURL: coderAccessURL, + coderAccessPort: coderAccessPort, aibridgeProviderFromHost: aibridgeProviderFromHost, caCert: certPEM, allowedPrivateRanges: allowedPrivateRanges, @@ -782,7 +784,7 @@ func (s *Server) isBlockedIP(ip net.IP, hostname string, port string) bool { // block connections to its own deployment. Hostname-based (not IP-based) // to handle dynamic IPs (DNS changes, load balancers, k8s rescheduling). // The port is normalized at startup to handle URLs without explicit ports. - if strings.EqualFold(hostname, s.coderAccessURL.Hostname()) && port == s.coderAccessURL.Port() { + if strings.EqualFold(hostname, s.coderAccessURL.Hostname()) && port == s.coderAccessPort { return false } diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go index a25cbb4c1e1..b3beea7ad2a 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go @@ -574,8 +574,7 @@ func TestNew(t *testing.T) { AIBridgeProviderFromHost: testProviderFromHost, }) require.NoError(t, err) - require.Equal(t, "localhost", srv.CoderAccessURL().Hostname()) - require.Equal(t, "80", srv.CoderAccessURL().Port()) + require.Equal(t, "localhost", srv.CoderAccessURL().Host) }) t.Run("CoderAccessURLDefaultHTTPSPort", func(t *testing.T) { @@ -593,8 +592,7 @@ func TestNew(t *testing.T) { AIBridgeProviderFromHost: testProviderFromHost, }) require.NoError(t, err) - require.Equal(t, "localhost", srv.CoderAccessURL().Hostname()) - require.Equal(t, "443", srv.CoderAccessURL().Port()) + require.Equal(t, "localhost", srv.CoderAccessURL().Host) }) t.Run("CoderAccessURLExplicitPort", func(t *testing.T) { @@ -1022,6 +1020,49 @@ func TestNew(t *testing.T) { require.NoError(t, err) require.NotNil(t, srv) }) + + t.Run("CoderAccessURLHostPreserved", func(t *testing.T) { + t.Parallel() + + mitmCertFile, mitmKeyFile := getSharedTestMITMCert(t) + logger := slogtest.Make(t, nil) + + srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{ + ListenAddr: "127.0.0.1:0", + CoderAccessURL: "https://coder.example.com", + MITMCertFile: mitmCertFile, + MITMKeyFile: mitmKeyFile, + DomainAllowlist: []string{aibridgeproxyd.HostAnthropic}, + AIBridgeProviderFromHost: testProviderFromHost, + AllowedPrivateCIDRs: []string{"127.0.0.1/32"}, + }) + require.NoError(t, err) + t.Cleanup(func() { _ = srv.Close() }) + + require.Equal(t, "coder.example.com", srv.CoderAccessURL().Host, + "Host must not have :443 appended") + }) + + t.Run("CoderAccessURLExplicitPortPreserved", func(t *testing.T) { + t.Parallel() + + mitmCertFile, mitmKeyFile := getSharedTestMITMCert(t) + logger := slogtest.Make(t, nil) + + srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{ + ListenAddr: "127.0.0.1:0", + CoderAccessURL: "https://coder.example.com:8443", + MITMCertFile: mitmCertFile, + MITMKeyFile: mitmKeyFile, + DomainAllowlist: []string{aibridgeproxyd.HostAnthropic}, + AIBridgeProviderFromHost: testProviderFromHost, + AllowedPrivateCIDRs: []string{"127.0.0.1/32"}, + }) + require.NoError(t, err) + t.Cleanup(func() { _ = srv.Close() }) + + require.Equal(t, "coder.example.com:8443", srv.CoderAccessURL().Host) + }) } func TestClose(t *testing.T) { diff --git a/scaletest/agentconn/run.go b/scaletest/agentconn/run.go index 4a4587e478d..f26db4355a6 100644 --- a/scaletest/agentconn/run.go +++ b/scaletest/agentconn/run.go @@ -214,7 +214,7 @@ func verifyConnection(ctx context.Context, logs io.Writer, conn workspacesdk.Age ctx, span := tracing.StartSpan(ctx) defer span.End() - client := agentHTTPClient(conn) + client := conn.AppHTTPClient() for i := 0; i < verifyConnectionAttempts; i++ { _, _ = fmt.Fprintf(logs, "\tVerify connection attempt %d/%d...\n", i+1, verifyConnectionAttempts) verifyCtx, cancel := context.WithTimeout(ctx, defaultRequestTimeout) @@ -258,7 +258,7 @@ func performInitialConnections(ctx context.Context, logs io.Writer, conn workspa defer span.End() _, _ = fmt.Fprintln(logs, "Performing initial service connections...") - client := agentHTTPClient(conn) + client := conn.AppHTTPClient() for i, connSpec := range specs { _, _ = fmt.Fprintf(logs, "\t%d. %s\n", i, connSpec.URL) @@ -292,7 +292,7 @@ func holdConnection(ctx context.Context, logs io.Writer, conn workspacesdk.Agent defer span.End() eg, egCtx := errgroup.WithContext(ctx) - client := agentHTTPClient(conn) + client := conn.AppHTTPClient() if len(specs) > 0 { _, _ = fmt.Fprintln(logs, "\nStarting connection loops...") } @@ -362,26 +362,3 @@ func holdConnection(ctx context.Context, logs io.Writer, conn workspacesdk.Agent return nil } - -func agentHTTPClient(conn workspacesdk.AgentConn) *http.Client { - return &http.Client{ - Transport: &http.Transport{ - DisableKeepAlives: true, - DialContext: func(ctx context.Context, _ string, addr string) (net.Conn, error) { - _, port, err := net.SplitHostPort(addr) - if err != nil { - return nil, xerrors.Errorf("split host port %q: %w", addr, err) - } - - portUint, err := strconv.ParseUint(port, 10, 16) - if err != nil { - return nil, xerrors.Errorf("parse port %q: %w", port, err) - } - - // Addr doesn't matter here, besides the port. DialContext will - // automatically choose the right IP to dial. - return conn.DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", portUint)) - }, - }, - } -} diff --git a/site/src/hooks/useExternalAuth.ts b/site/src/hooks/useExternalAuth.ts index 6c0db550d89..81dcae5de06 100644 --- a/site/src/hooks/useExternalAuth.ts +++ b/site/src/hooks/useExternalAuth.ts @@ -5,13 +5,16 @@ import { templateVersionExternalAuth } from "#/api/queries/templates"; export type ExternalAuthPollingState = "idle" | "polling" | "abandoned"; export const useExternalAuth = (versionId: string | undefined) => { - const [externalAuthPollingState, setExternalAuthPollingState] = - useState("idle"); + const [pollingState, setPollingState] = useState< + Record + >({}); - const startPollingExternalAuth = useCallback(() => { - setExternalAuthPollingState("polling"); + const startPollingExternalAuth = useCallback((providerId: string) => { + setPollingState((prev) => ({ ...prev, [providerId]: "polling" })); }, []); + const isAnyPolling = Object.values(pollingState).some((s) => s === "polling"); + const { data: externalAuth, isPending: isLoadingExternalAuth, @@ -19,37 +22,58 @@ export const useExternalAuth = (versionId: string | undefined) => { } = useQuery({ ...templateVersionExternalAuth(versionId ?? ""), enabled: Boolean(versionId), - refetchInterval: externalAuthPollingState === "polling" ? 1000 : false, + refetchInterval: isAnyPolling ? 1000 : false, }); - const allSignedIn = externalAuth?.every((it) => it.authenticated); - + // Stop polling individual providers once they authenticate. useEffect(() => { - if (allSignedIn) { - setExternalAuthPollingState("idle"); + if (!externalAuth) { return; } + setPollingState((prev) => { + let changed = false; + const next = { ...prev }; + for (const auth of externalAuth) { + if (auth.authenticated && next[auth.id] === "polling") { + next[auth.id] = "idle"; + changed = true; + } + } + return changed ? next : prev; + }); + }, [externalAuth]); - if (externalAuthPollingState !== "polling") { + // Per-provider 60-second timeout. + useEffect(() => { + const pollingIds = Object.entries(pollingState) + .filter(([, authPollingState]) => authPollingState === "polling") + .map(([id]) => id); + + if (pollingIds.length === 0) { return; } - // Poll for a maximum of one minute - const quitPolling = setTimeout( - () => setExternalAuthPollingState("abandoned"), - 60_000, + const timers = pollingIds.map((id) => + setTimeout(() => { + setPollingState((prev) => + prev[id] === "polling" ? { ...prev, [id]: "abandoned" } : prev, + ); + }, 60_000), ); + return () => { - clearTimeout(quitPolling); + for (const t of timers) { + clearTimeout(t); + } }; - }, [externalAuthPollingState, allSignedIn]); + }, [pollingState]); return { startPollingExternalAuth, externalAuth, - externalAuthPollingState, + externalAuthPollingState: pollingState, isLoadingExternalAuth, externalAuthError: error, - isPollingExternalAuth: externalAuthPollingState === "polling", + isPollingExternalAuth: isAnyPolling, }; }; diff --git a/site/src/modules/tasks/TaskPrompt/TaskPrompt.stories.tsx b/site/src/modules/tasks/TaskPrompt/TaskPrompt.stories.tsx index 43e0b6083f4..809c9256ebd 100644 --- a/site/src/modules/tasks/TaskPrompt/TaskPrompt.stories.tsx +++ b/site/src/modules/tasks/TaskPrompt/TaskPrompt.stories.tsx @@ -10,6 +10,7 @@ import { MockTasks, MockTemplate, MockTemplateVersion, + MockTemplateVersionExternalAuthAzure, MockTemplateVersionExternalAuthGithub, MockTemplateVersionExternalAuthGithubAuthenticated, MockUserOwner, @@ -387,6 +388,39 @@ export const MissingExternalAuth: Story = { }, }; +export const MissingExternalAuthMultipleProviders: Story = { + beforeEach: () => { + spyOn(API, "getTasks") + .mockResolvedValueOnce(MockTasks) + .mockResolvedValue([MockNewTaskData, ...MockTasks]); + spyOn(API, "createTask").mockResolvedValue(MockTask); + spyOn(API, "getTemplateVersionExternalAuth").mockResolvedValue([ + MockTemplateVersionExternalAuthGithub, + MockTemplateVersionExternalAuthAzure, + ]); + // Prevent the auth button from actually opening a popup. + spyOn(window, "open").mockReturnValue(null); + }, + play: async ({ canvasElement, step }) => { + const canvas = within(canvasElement); + + const githubButton = await canvas.findByRole("button", { + name: /connect to github/i, + }); + const azureButton = await canvas.findByRole("button", { + name: /connect to azure/i, + }); + + await step("Click GitHub auth button", async () => { + await userEvent.click(githubButton); + }); + + await step("Azure button remains enabled", () => { + expect(azureButton).toBeEnabled(); + }); + }, +}; + export const ExternalAuthError: Story = { beforeEach: () => { spyOn(API, "getTasks") diff --git a/site/src/modules/tasks/TaskPrompt/TaskPrompt.tsx b/site/src/modules/tasks/TaskPrompt/TaskPrompt.tsx index 8468f755816..21481e6c2d4 100644 --- a/site/src/modules/tasks/TaskPrompt/TaskPrompt.tsx +++ b/site/src/modules/tasks/TaskPrompt/TaskPrompt.tsx @@ -436,14 +436,14 @@ const ExternalAuthButtons: FC = ({ versionId, missedExternalAuth, }) => { - const { - startPollingExternalAuth, - isPollingExternalAuth, - externalAuthPollingState, - } = useExternalAuth(versionId); - const shouldRetry = externalAuthPollingState === "abandoned"; + const { startPollingExternalAuth, externalAuthPollingState } = + useExternalAuth(versionId); return missedExternalAuth.map((auth) => { + const isPollingExternalAuth = + externalAuthPollingState[auth.id] === "polling"; + const shouldRetry = externalAuthPollingState[auth.id] === "abandoned"; + return (
diff --git a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.stories.tsx b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.stories.tsx index a4f2f15ee6a..d131c874143 100644 --- a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.stories.tsx +++ b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.stories.tsx @@ -31,7 +31,8 @@ const meta: Meta = { activePath: "main.tf", template: MockTemplate, templateVersion: MockTemplateVersion, - defaultFileTree: MockTemplateVersionFileTree, + fileTree: MockTemplateVersionFileTree, + onFileTreeChange: action("onFileTreeChange"), onPublish: action("onPublish"), onConfirmPublish: action("onConfirmPublish"), onCancelPublish: action("onCancelPublish"), diff --git a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.tsx b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.tsx index 5393f496e32..1af98fae4a6 100644 --- a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.tsx +++ b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditor.tsx @@ -78,7 +78,8 @@ type Tab = "logs" | "resources" | undefined; // Undefined is to hide the tab interface TemplateVersionEditorProps { template: Template; templateVersion: TemplateVersion; - defaultFileTree: FileTree; + fileTree: FileTree; + onFileTreeChange: (updater: (fileTree: FileTree) => FileTree) => void; buildLogs?: ProvisionerJobLog[]; resources?: WorkspaceResource[]; isBuilding: boolean; @@ -108,7 +109,8 @@ export const TemplateVersionEditor: FC = ({ canPublish, template, templateVersion, - defaultFileTree, + fileTree, + onFileTreeChange, onPreview, onPublish, onConfirmPublish, @@ -133,7 +135,6 @@ export const TemplateVersionEditor: FC = ({ const navigate = useNavigate(); const getLink = useLinks(); const [selectedTab, setSelectedTab] = useState(defaultTab); - const [fileTree, setFileTree] = useState(defaultFileTree); const [createFileOpen, setCreateFileOpen] = useState(false); const [deleteFileOpen, setDeleteFileOpen] = useState(); const [renameFileOpen, setRenameFileOpen] = useState(); @@ -349,7 +350,9 @@ export const TemplateVersionEditor: FC = ({ }} checkExists={(path) => existsFile(path, fileTree)} onConfirm={(path) => { - setFileTree((fileTree) => createFile(path, fileTree, "")); + onFileTreeChange((fileTree) => + createFile(path, fileTree, ""), + ); onActivePathChange(path); setCreateFileOpen(false); setDirty(true); @@ -360,7 +363,7 @@ export const TemplateVersionEditor: FC = ({ if (!deleteFileOpen) { throw new Error("delete file must be set"); } - setFileTree((fileTree) => + onFileTreeChange((fileTree) => removeFile(deleteFileOpen, fileTree), ); setDeleteFileOpen(undefined); @@ -385,7 +388,7 @@ export const TemplateVersionEditor: FC = ({ if (!renameFileOpen) { return; } - setFileTree((fileTree) => + onFileTreeChange((fileTree) => moveFile(renameFileOpen, newPath, fileTree), ); onActivePathChange(newPath); @@ -431,7 +434,7 @@ export const TemplateVersionEditor: FC = ({ if (!activePath) { return; } - setFileTree((fileTree) => + onFileTreeChange((fileTree) => updateFile(activePath, value, fileTree), ); setDirty(true); diff --git a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.test.tsx b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.test.tsx index 157ec23a259..2c660fcbaa1 100644 --- a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.test.tsx +++ b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.test.tsx @@ -295,6 +295,78 @@ test("Preserves the currently open file path when building a template version", expect(router.state.location.search).toBe("?path=myfile.tf"); }); +test("Creating a new file opens it in the editor", async () => { + const user = userEvent.setup(); + const { router } = renderWithAuth(, { + route: `/templates/${MockTemplate.name}/versions/${MockTemplateVersion.name}/edit`, + path: "/templates/:template/versions/:version/edit", + }); + + // Wait for the default entrypoint file to load. + const editor = await screen.findByTestId("monaco-editor"); + await waitFor(() => { + expect(editor).not.toHaveValue(""); + }); + + const createButton = await screen.findByRole("button", { + name: "Create File", + }); + await user.click(createButton); + + const dialog = await screen.findByTestId("dialog"); + const pathField = within(dialog).getByLabelText("File Path"); + await user.type(pathField, "newfile.tf"); + await user.click(within(dialog).getByRole("button", { name: "Create" })); + + // The new (empty) file should be opened in the editor and the URL path + // query parameter should reflect the new file. + await waitFor(() => { + expect(screen.getByTestId("monaco-editor")).toHaveValue(""); + }); + expect(router.state.location.search).toBe("?path=newfile.tf"); +}); + +test("Renaming a file does not throw and opens the new path", async () => { + const user = userEvent.setup(); + const { router } = renderWithAuth(, { + route: `/templates/${MockTemplate.name}/versions/${MockTemplateVersion.name}/edit`, + path: "/templates/:template/versions/:version/edit", + }); + + // Wait for the default entrypoint file to load and capture its content so + // we can confirm the same content is shown after renaming. + const editor = await screen.findByTestId("monaco-editor"); + await waitFor(() => { + expect(editor).not.toHaveValue(""); + }); + if (!(editor instanceof HTMLTextAreaElement)) { + throw new Error("editor is not a textarea"); + } + const originalContent = editor.value; + + // Open the file actions menu for the active file and click Rename. + const fileActions = await screen.findByRole("button", { + name: "File actions", + }); + await user.click(fileActions); + await user.click(await screen.findByRole("menuitem", { name: /rename/i })); + + const dialog = await screen.findByTestId("dialog"); + const pathField = within(dialog).getByLabelText("File Path"); + await user.clear(pathField); + await user.type(pathField, "renamed.tf"); + await user.click(within(dialog).getByRole("button", { name: "Rename" })); + + // The renamed file should still be open with its original content and + // the URL path query parameter should reflect the new name. Previously + // this path threw "File is not a text file" because the parent's stale + // file tree fell back to the old entrypoint name. + await waitFor(() => { + expect(screen.getByTestId("monaco-editor")).toHaveValue(originalContent); + }); + expect(router.state.location.search).toBe("?path=renamed.tf"); +}); + describe.each([ { testName: "Do not ask when template version has no errors", diff --git a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.tsx b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.tsx index 2f528263326..1d95b12e784 100644 --- a/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.tsx +++ b/site/src/pages/TemplateVersionEditorPage/TemplateVersionEditorPage.tsx @@ -72,7 +72,7 @@ const TemplateVersionEditorPage: FC = () => { const logs = useWatchVersionLogs(activeTemplateVersion, { onDone: activeTemplateVersionQuery.refetch, }); - const { fileTree, tarFile } = useFileTree(activeTemplateVersion); + const { fileTree, setFileTree, tarFile } = useFileTree(activeTemplateVersion); const { missingVariables, setIsMissingVariablesDialogOpen, @@ -138,7 +138,10 @@ const TemplateVersionEditorPage: FC = () => { onActivePathChange={onActivePathChange} template={templateQuery.data} templateVersion={activeTemplateVersion} - defaultFileTree={fileTree} + fileTree={fileTree} + onFileTreeChange={(updater) => { + setFileTree((current) => (current ? updater(current) : current)); + }} onPreview={async (newFileTree) => { if (!tarFile) { return; @@ -246,13 +249,8 @@ const useFileTree = (templateVersion: TemplateVersion | undefined) => { ...file(templateVersion?.job.file_id ?? ""), enabled: templateVersion !== undefined, }); - const [state, setState] = useState<{ - fileTree?: FileTree; - tarFile?: TarReader; - }>({ - fileTree: undefined, - tarFile: undefined, - }); + const [fileTree, setFileTree] = useState(undefined); + const [tarFile, setTarFile] = useState(undefined); useEffect(() => { let stale = false; @@ -264,8 +262,8 @@ const useFileTree = (templateVersion: TemplateVersion | undefined) => { if (stale) { return; } - const fileTree = createTemplateVersionFileTree(tarFile); - setState({ fileTree, tarFile }); + setFileTree(createTemplateVersionFileTree(tarFile)); + setTarFile(tarFile); } catch (error) { console.error(error); toast.error("Error on initializing the editor.", { @@ -283,7 +281,7 @@ const useFileTree = (templateVersion: TemplateVersion | undefined) => { }; }, [fileQuery.data]); - return state; + return { fileTree, setFileTree, tarFile }; }; const useMissingVariables = (templateVersion: TemplateVersion | undefined) => { diff --git a/site/src/testHelpers/entities.ts b/site/src/testHelpers/entities.ts index d5ea25ae39c..71e5d38f3f6 100644 --- a/site/src/testHelpers/entities.ts +++ b/site/src/testHelpers/entities.ts @@ -3518,6 +3518,16 @@ export const MockTemplateVersionExternalAuthGithubAuthenticated: TypesGen.Templa display_name: "GitHub", }; +export const MockTemplateVersionExternalAuthAzure: TypesGen.TemplateVersionExternalAuth = + { + id: "azure", + type: "azure", + authenticate_url: "https://example.com/external-auth/azure", + authenticated: false, + display_icon: "/icon/azure.svg", + display_name: "Azure", + }; + export const MockDeploymentStats: TypesGen.DeploymentStats = { aggregated_from: "2023-03-06T19:08:55.211625Z", collected_at: "2023-03-06T19:12:55.211625Z",