mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(site): use websocket connection for devcontainer updates (#18808)
Instead of polling every 10 seconds, we instead use a WebSocket connection for more timely updates.
This commit is contained in:
@@ -2,8 +2,10 @@ package agentcontainers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"maps"
|
||||
"net/http"
|
||||
"os"
|
||||
"path"
|
||||
@@ -30,6 +32,7 @@ import (
|
||||
"github.com/coder/coder/v2/codersdk/agentsdk"
|
||||
"github.com/coder/coder/v2/provisioner"
|
||||
"github.com/coder/quartz"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -74,6 +77,7 @@ type API struct {
|
||||
|
||||
mu sync.RWMutex // Protects the following fields.
|
||||
initDone chan struct{} // Closed by Init.
|
||||
updateChans []chan struct{}
|
||||
closed bool
|
||||
containers codersdk.WorkspaceAgentListContainersResponse // Output from the last list operation.
|
||||
containersErr error // Error from the last list operation.
|
||||
@@ -535,6 +539,7 @@ func (api *API) Routes() http.Handler {
|
||||
r.Use(ensureInitDoneMW)
|
||||
|
||||
r.Get("/", api.handleList)
|
||||
r.Get("/watch", api.watchContainers)
|
||||
// TODO(mafredri): Simplify this route as the previous /devcontainers
|
||||
// /-route was dropped. We can drop the /devcontainers prefix here too.
|
||||
r.Route("/devcontainers/{devcontainer}", func(r chi.Router) {
|
||||
@@ -544,6 +549,88 @@ func (api *API) Routes() http.Handler {
|
||||
return r
|
||||
}
|
||||
|
||||
func (api *API) broadcastUpdatesLocked() {
|
||||
// Broadcast state changes to WebSocket listeners.
|
||||
for _, ch := range api.updateChans {
|
||||
select {
|
||||
case ch <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (api *API) watchContainers(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
conn, err := websocket.Accept(rw, r, nil)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to upgrade connection to websocket.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Here we close the websocket for reading, so that the websocket library will handle pings and
|
||||
// close frames.
|
||||
_ = conn.CloseRead(context.Background())
|
||||
|
||||
ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageText)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
go httpapi.Heartbeat(ctx, conn)
|
||||
|
||||
updateCh := make(chan struct{}, 1)
|
||||
|
||||
api.mu.Lock()
|
||||
api.updateChans = append(api.updateChans, updateCh)
|
||||
api.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
api.mu.Lock()
|
||||
api.updateChans = slices.DeleteFunc(api.updateChans, func(ch chan struct{}) bool {
|
||||
return ch == updateCh
|
||||
})
|
||||
close(updateCh)
|
||||
api.mu.Unlock()
|
||||
}()
|
||||
|
||||
encoder := json.NewEncoder(wsNetConn)
|
||||
|
||||
ct, err := api.getContainers()
|
||||
if err != nil {
|
||||
api.logger.Error(ctx, "unable to get containers", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
if err := encoder.Encode(ct); err != nil {
|
||||
api.logger.Error(ctx, "encode container list", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-api.ctx.Done():
|
||||
return
|
||||
|
||||
case <-ctx.Done():
|
||||
return
|
||||
|
||||
case <-updateCh:
|
||||
ct, err := api.getContainers()
|
||||
if err != nil {
|
||||
api.logger.Error(ctx, "unable to get containers", slog.Error(err))
|
||||
continue
|
||||
}
|
||||
|
||||
if err := encoder.Encode(ct); err != nil {
|
||||
api.logger.Error(ctx, "encode container list", slog.Error(err))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// handleList handles the HTTP request to list containers.
|
||||
func (api *API) handleList(rw http.ResponseWriter, r *http.Request) {
|
||||
ct, err := api.getContainers()
|
||||
@@ -583,8 +670,26 @@ func (api *API) updateContainers(ctx context.Context) error {
|
||||
api.mu.Lock()
|
||||
defer api.mu.Unlock()
|
||||
|
||||
var previouslyKnownDevcontainers map[string]codersdk.WorkspaceAgentDevcontainer
|
||||
if len(api.updateChans) > 0 {
|
||||
previouslyKnownDevcontainers = maps.Clone(api.knownDevcontainers)
|
||||
}
|
||||
|
||||
api.processUpdatedContainersLocked(ctx, updated)
|
||||
|
||||
if len(api.updateChans) > 0 {
|
||||
statesAreEqual := maps.EqualFunc(
|
||||
previouslyKnownDevcontainers,
|
||||
api.knownDevcontainers,
|
||||
func(dc1, dc2 codersdk.WorkspaceAgentDevcontainer) bool {
|
||||
return dc1.Equals(dc2)
|
||||
})
|
||||
|
||||
if !statesAreEqual {
|
||||
api.broadcastUpdatesLocked()
|
||||
}
|
||||
}
|
||||
|
||||
api.logger.Debug(ctx, "containers updated successfully", slog.F("container_count", len(api.containers.Containers)), slog.F("warning_count", len(api.containers.Warnings)), slog.F("devcontainer_count", len(api.knownDevcontainers)))
|
||||
|
||||
return nil
|
||||
@@ -955,6 +1060,8 @@ func (api *API) handleDevcontainerRecreate(w http.ResponseWriter, r *http.Reques
|
||||
dc.Container = nil
|
||||
dc.Error = ""
|
||||
api.knownDevcontainers[dc.WorkspaceFolder] = dc
|
||||
api.broadcastUpdatesLocked()
|
||||
|
||||
go func() {
|
||||
_ = api.CreateDevcontainer(dc.WorkspaceFolder, dc.ConfigPath, WithRemoveExistingContainer())
|
||||
}()
|
||||
@@ -1070,6 +1177,7 @@ func (api *API) CreateDevcontainer(workspaceFolder, configPath string, opts ...D
|
||||
dc.Error = ""
|
||||
api.recreateSuccessTimes[dc.WorkspaceFolder] = api.clock.Now("agentcontainers", "recreate", "successTimes")
|
||||
api.knownDevcontainers[dc.WorkspaceFolder] = dc
|
||||
api.broadcastUpdatesLocked()
|
||||
api.mu.Unlock()
|
||||
|
||||
// Ensure an immediate refresh to accurately reflect the
|
||||
|
||||
@@ -36,6 +36,7 @@ import (
|
||||
"github.com/coder/coder/v2/pty"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
// fakeContainerCLI implements the agentcontainers.ContainerCLI interface for
|
||||
@@ -441,6 +442,178 @@ func TestAPI(t *testing.T) {
|
||||
logbuf.Reset()
|
||||
})
|
||||
|
||||
t.Run("Watch", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fakeContainer1 := fakeContainer(t, func(c *codersdk.WorkspaceAgentContainer) {
|
||||
c.ID = "container1"
|
||||
c.FriendlyName = "devcontainer1"
|
||||
c.Image = "busybox:latest"
|
||||
c.Labels = map[string]string{
|
||||
agentcontainers.DevcontainerLocalFolderLabel: "/home/coder/project1",
|
||||
agentcontainers.DevcontainerConfigFileLabel: "/home/coder/project1/.devcontainer/devcontainer.json",
|
||||
}
|
||||
})
|
||||
|
||||
fakeContainer2 := fakeContainer(t, func(c *codersdk.WorkspaceAgentContainer) {
|
||||
c.ID = "container2"
|
||||
c.FriendlyName = "devcontainer2"
|
||||
c.Image = "ubuntu:latest"
|
||||
c.Labels = map[string]string{
|
||||
agentcontainers.DevcontainerLocalFolderLabel: "/home/coder/project2",
|
||||
agentcontainers.DevcontainerConfigFileLabel: "/home/coder/project2/.devcontainer/devcontainer.json",
|
||||
}
|
||||
})
|
||||
|
||||
stages := []struct {
|
||||
containers []codersdk.WorkspaceAgentContainer
|
||||
expected codersdk.WorkspaceAgentListContainersResponse
|
||||
}{
|
||||
{
|
||||
containers: []codersdk.WorkspaceAgentContainer{fakeContainer1},
|
||||
expected: codersdk.WorkspaceAgentListContainersResponse{
|
||||
Containers: []codersdk.WorkspaceAgentContainer{fakeContainer1},
|
||||
Devcontainers: []codersdk.WorkspaceAgentDevcontainer{
|
||||
{
|
||||
Name: "project1",
|
||||
WorkspaceFolder: fakeContainer1.Labels[agentcontainers.DevcontainerLocalFolderLabel],
|
||||
ConfigPath: fakeContainer1.Labels[agentcontainers.DevcontainerConfigFileLabel],
|
||||
Status: "running",
|
||||
Container: &fakeContainer1,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
containers: []codersdk.WorkspaceAgentContainer{fakeContainer1, fakeContainer2},
|
||||
expected: codersdk.WorkspaceAgentListContainersResponse{
|
||||
Containers: []codersdk.WorkspaceAgentContainer{fakeContainer1, fakeContainer2},
|
||||
Devcontainers: []codersdk.WorkspaceAgentDevcontainer{
|
||||
{
|
||||
Name: "project1",
|
||||
WorkspaceFolder: fakeContainer1.Labels[agentcontainers.DevcontainerLocalFolderLabel],
|
||||
ConfigPath: fakeContainer1.Labels[agentcontainers.DevcontainerConfigFileLabel],
|
||||
Status: "running",
|
||||
Container: &fakeContainer1,
|
||||
},
|
||||
{
|
||||
Name: "project2",
|
||||
WorkspaceFolder: fakeContainer2.Labels[agentcontainers.DevcontainerLocalFolderLabel],
|
||||
ConfigPath: fakeContainer2.Labels[agentcontainers.DevcontainerConfigFileLabel],
|
||||
Status: "running",
|
||||
Container: &fakeContainer2,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
containers: []codersdk.WorkspaceAgentContainer{fakeContainer2},
|
||||
expected: codersdk.WorkspaceAgentListContainersResponse{
|
||||
Containers: []codersdk.WorkspaceAgentContainer{fakeContainer2},
|
||||
Devcontainers: []codersdk.WorkspaceAgentDevcontainer{
|
||||
{
|
||||
Name: "",
|
||||
WorkspaceFolder: fakeContainer1.Labels[agentcontainers.DevcontainerLocalFolderLabel],
|
||||
ConfigPath: fakeContainer1.Labels[agentcontainers.DevcontainerConfigFileLabel],
|
||||
Status: "stopped",
|
||||
Container: nil,
|
||||
},
|
||||
{
|
||||
Name: "project2",
|
||||
WorkspaceFolder: fakeContainer2.Labels[agentcontainers.DevcontainerLocalFolderLabel],
|
||||
ConfigPath: fakeContainer2.Labels[agentcontainers.DevcontainerConfigFileLabel],
|
||||
Status: "running",
|
||||
Container: &fakeContainer2,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
var (
|
||||
ctx = testutil.Context(t, testutil.WaitShort)
|
||||
mClock = quartz.NewMock(t)
|
||||
updaterTickerTrap = mClock.Trap().TickerFunc("updaterLoop")
|
||||
mCtrl = gomock.NewController(t)
|
||||
mLister = acmock.NewMockContainerCLI(mCtrl)
|
||||
logger = slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
|
||||
)
|
||||
|
||||
// Set up initial state for immediate send on connection
|
||||
mLister.EXPECT().List(gomock.Any()).Return(codersdk.WorkspaceAgentListContainersResponse{Containers: stages[0].containers}, nil)
|
||||
mLister.EXPECT().DetectArchitecture(gomock.Any(), gomock.Any()).Return("<none>", nil).AnyTimes()
|
||||
|
||||
api := agentcontainers.NewAPI(logger,
|
||||
agentcontainers.WithClock(mClock),
|
||||
agentcontainers.WithContainerCLI(mLister),
|
||||
agentcontainers.WithWatcher(watcher.NewNoop()),
|
||||
)
|
||||
api.Start()
|
||||
defer api.Close()
|
||||
|
||||
srv := httptest.NewServer(api.Routes())
|
||||
defer srv.Close()
|
||||
|
||||
updaterTickerTrap.MustWait(ctx).MustRelease(ctx)
|
||||
defer updaterTickerTrap.Close()
|
||||
|
||||
client, res, err := websocket.Dial(ctx, srv.URL+"/watch", nil)
|
||||
require.NoError(t, err)
|
||||
if res != nil && res.Body != nil {
|
||||
defer res.Body.Close()
|
||||
}
|
||||
|
||||
// Read initial state sent immediately on connection
|
||||
mt, msg, err := client.Read(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, websocket.MessageText, mt)
|
||||
|
||||
var got codersdk.WorkspaceAgentListContainersResponse
|
||||
err = json.Unmarshal(msg, &got)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, stages[0].expected.Containers, got.Containers)
|
||||
require.Len(t, got.Devcontainers, len(stages[0].expected.Devcontainers))
|
||||
for j, expectedDev := range stages[0].expected.Devcontainers {
|
||||
gotDev := got.Devcontainers[j]
|
||||
require.Equal(t, expectedDev.Name, gotDev.Name)
|
||||
require.Equal(t, expectedDev.WorkspaceFolder, gotDev.WorkspaceFolder)
|
||||
require.Equal(t, expectedDev.ConfigPath, gotDev.ConfigPath)
|
||||
require.Equal(t, expectedDev.Status, gotDev.Status)
|
||||
require.Equal(t, expectedDev.Container, gotDev.Container)
|
||||
}
|
||||
|
||||
// Process remaining stages through updater loop
|
||||
for i, stage := range stages[1:] {
|
||||
mLister.EXPECT().List(gomock.Any()).Return(codersdk.WorkspaceAgentListContainersResponse{Containers: stage.containers}, nil)
|
||||
|
||||
// Given: We allow the update loop to progress
|
||||
_, aw := mClock.AdvanceNext()
|
||||
aw.MustWait(ctx)
|
||||
|
||||
// When: We attempt to read a message from the socket.
|
||||
mt, msg, err := client.Read(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, websocket.MessageText, mt)
|
||||
|
||||
// Then: We expect the receieved message matches the expected response.
|
||||
var got codersdk.WorkspaceAgentListContainersResponse
|
||||
err = json.Unmarshal(msg, &got)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, stages[i+1].expected.Containers, got.Containers)
|
||||
require.Len(t, got.Devcontainers, len(stages[i+1].expected.Devcontainers))
|
||||
for j, expectedDev := range stages[i+1].expected.Devcontainers {
|
||||
gotDev := got.Devcontainers[j]
|
||||
require.Equal(t, expectedDev.Name, gotDev.Name)
|
||||
require.Equal(t, expectedDev.WorkspaceFolder, gotDev.WorkspaceFolder)
|
||||
require.Equal(t, expectedDev.ConfigPath, gotDev.ConfigPath)
|
||||
require.Equal(t, expectedDev.Status, gotDev.Status)
|
||||
require.Equal(t, expectedDev.Container, gotDev.Container)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// List tests the API.getContainers method using a mock
|
||||
// implementation. It specifically tests caching behavior.
|
||||
t.Run("List", func(t *testing.T) {
|
||||
|
||||
Reference in New Issue
Block a user