mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: agents git watch backend (#22565)
Adds real-time git status watching for workspace agents, so the frontend
can subscribe over WebSocket and show
git file changes in near real-time.
1. Subscription is scoped to a **chat** via `GET
/api/experimental/chats/{chat}/git/watch`.
2. The workspace agent automatically determines which paths to watch
based on tool calls made by the chat (and its ancestor chats).
3. Workspace agent polls subscribed repo working trees on a 30s
interval, on tools calls, and on explicit `refresh` from the client.
4. Scans are rate-limited to at most once per second.
5. Edited paths are tracked **in-memory** inside the workspace agent.
There is no database persistence — state is lost on agent restart. This
will be addresses in a future PR.
6. Messages sent over WebSocket include a full-repo snapshot (unified
diff, branch, origin). A new message is emitted only when the snapshot
changes.
This PR was implemented with AI with me closely controlling what it's
doing. The code follows a plan file that was updated continuously during
implementation. Here's the file if you'd like to see it:
[project.md](https://gist.github.com/hugodutka/8722cf80c92f8a56555f7bc595b770e2).
It reflects the current state of the PR.
This commit is contained in:
Generated
+29
@@ -495,6 +495,35 @@ const docTemplate = `{
|
||||
}
|
||||
}
|
||||
},
|
||||
"/chats/{chat}/git/watch": {
|
||||
"get": {
|
||||
"security": [
|
||||
{
|
||||
"CoderSessionToken": []
|
||||
}
|
||||
],
|
||||
"tags": [
|
||||
"Chats"
|
||||
],
|
||||
"summary": "Watch git changes for a chat.",
|
||||
"operationId": "watch-chat-git",
|
||||
"parameters": [
|
||||
{
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "Chat ID",
|
||||
"name": "chat",
|
||||
"in": "path",
|
||||
"required": true
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"101": {
|
||||
"description": "Switching Protocols"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/chats/{chat}/unarchive": {
|
||||
"post": {
|
||||
"tags": [
|
||||
|
||||
Generated
+27
@@ -422,6 +422,33 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"/chats/{chat}/git/watch": {
|
||||
"get": {
|
||||
"security": [
|
||||
{
|
||||
"CoderSessionToken": []
|
||||
}
|
||||
],
|
||||
"tags": ["Chats"],
|
||||
"summary": "Watch git changes for a chat.",
|
||||
"operationId": "watch-chat-git",
|
||||
"parameters": [
|
||||
{
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "Chat ID",
|
||||
"name": "chat",
|
||||
"in": "path",
|
||||
"required": true
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"101": {
|
||||
"description": "Switching Protocols"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/chats/{chat}/unarchive": {
|
||||
"post": {
|
||||
"tags": ["Chats"],
|
||||
|
||||
@@ -2009,6 +2009,23 @@ func (p *Server) runChat(
|
||||
conn = agentConn
|
||||
releaseConn = agentRelease
|
||||
chatStateMu.Unlock()
|
||||
|
||||
// Inject chat identity headers so agent-side
|
||||
// handlers can track which paths this chat edits.
|
||||
var ancestorIDs []string
|
||||
if chatSnapshot.ParentChatID.Valid {
|
||||
ancestorIDs = append(ancestorIDs, chatSnapshot.ParentChatID.UUID.String())
|
||||
}
|
||||
ancestorJSON, err := json.Marshal(ancestorIDs)
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "failed to marshal ancestor chat IDs", slog.Error(err))
|
||||
ancestorJSON = []byte("[]")
|
||||
}
|
||||
agentConn.SetExtraHeaders(http.Header{
|
||||
workspacesdk.CoderChatIDHeader: {chatSnapshot.ID.String()},
|
||||
workspacesdk.CoderAncestorChatIDsHeader: {string(ancestorJSON)},
|
||||
})
|
||||
|
||||
return agentConn, nil
|
||||
}
|
||||
currentConn := conn
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -16,6 +17,8 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/agent/agenttest"
|
||||
@@ -29,6 +32,8 @@ import (
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk/agentconnmock"
|
||||
"github.com/coder/coder/v2/provisioner/echo"
|
||||
proto "github.com/coder/coder/v2/provisionersdk/proto"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -1662,3 +1667,260 @@ func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T)
|
||||
!fromDB.LastError.Valid
|
||||
}, testutil.WaitMedium, testutil.IntervalFast)
|
||||
}
|
||||
|
||||
func TestHeaderInjection(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// seedWorkspaceAgent creates the DB entities needed so that
|
||||
// GetWorkspaceAgentsInLatestBuildByWorkspaceID returns an
|
||||
// agent for the given workspace.
|
||||
seedWorkspaceAgent := func(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
ps dbpubsub.Pubsub,
|
||||
ownerID uuid.UUID,
|
||||
orgID uuid.UUID,
|
||||
) (workspaceID uuid.UUID, agentID uuid.UUID) {
|
||||
t.Helper()
|
||||
|
||||
// TemplateVersion needs its own provisioner job.
|
||||
versionJob := dbgen.ProvisionerJob(t, db, ps, database.ProvisionerJob{
|
||||
OrganizationID: orgID,
|
||||
InitiatorID: ownerID,
|
||||
Type: database.ProvisionerJobTypeTemplateVersionImport,
|
||||
})
|
||||
tv := dbgen.TemplateVersion(t, db, database.TemplateVersion{
|
||||
OrganizationID: orgID,
|
||||
CreatedBy: ownerID,
|
||||
JobID: versionJob.ID,
|
||||
})
|
||||
templ := dbgen.Template(t, db, database.Template{
|
||||
OrganizationID: orgID,
|
||||
CreatedBy: ownerID,
|
||||
ActiveVersionID: tv.ID,
|
||||
})
|
||||
ws := dbgen.Workspace(t, db, database.WorkspaceTable{
|
||||
OwnerID: ownerID,
|
||||
OrganizationID: orgID,
|
||||
TemplateID: templ.ID,
|
||||
})
|
||||
buildJob := dbgen.ProvisionerJob(t, db, ps, database.ProvisionerJob{
|
||||
OrganizationID: orgID,
|
||||
InitiatorID: ownerID,
|
||||
Type: database.ProvisionerJobTypeWorkspaceBuild,
|
||||
})
|
||||
build := dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
|
||||
WorkspaceID: ws.ID,
|
||||
JobID: buildJob.ID,
|
||||
BuildNumber: 1,
|
||||
InitiatorID: ownerID,
|
||||
TemplateVersionID: tv.ID,
|
||||
})
|
||||
resource := dbgen.WorkspaceResource(t, db, database.WorkspaceResource{
|
||||
JobID: build.JobID,
|
||||
})
|
||||
agent := dbgen.WorkspaceAgent(t, db, database.WorkspaceAgent{
|
||||
ResourceID: resource.ID,
|
||||
})
|
||||
return ws.ID, agent.ID
|
||||
}
|
||||
|
||||
t.Run("WithParentChat", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
org, err := db.GetDefaultOrganization(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
workspaceID, expectedAgentID := seedWorkspaceAgent(t, db, ps, user.ID, org.ID)
|
||||
|
||||
// Set up the mock OpenAI to return a simple text response
|
||||
// so the chat finishes cleanly.
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("done")...,
|
||||
)
|
||||
})
|
||||
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
|
||||
|
||||
// Wire up the mock agent connection so we can capture
|
||||
// the headers passed to SetExtraHeaders.
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
|
||||
var capturedHeaders http.Header
|
||||
headersCaptured := make(chan struct{})
|
||||
|
||||
// SetExtraHeaders is called once when the connection
|
||||
// is first established.
|
||||
mockConn.EXPECT().SetExtraHeaders(gomock.Any()).Do(func(h http.Header) {
|
||||
capturedHeaders = h
|
||||
close(headersCaptured)
|
||||
})
|
||||
// resolveInstructions calls LS to look for instruction
|
||||
// files; return an error so it skips gracefully.
|
||||
mockConn.EXPECT().LS(gomock.Any(), gomock.Any(), gomock.Any()).Return(
|
||||
workspacesdk.LSResponse{}, xerrors.New("not found"),
|
||||
).AnyTimes()
|
||||
// The connection is closed when the chat finishes.
|
||||
mockConn.EXPECT().Close().Return(nil).AnyTimes()
|
||||
|
||||
agentConnFn := func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, expectedAgentID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := chatd.New(chatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: uuid.New(),
|
||||
Pubsub: ps,
|
||||
AgentConn: agentConnFn,
|
||||
PendingChatAcquireInterval: 10 * time.Millisecond,
|
||||
InFlightChatStaleAfter: testutil.WaitSuperLong,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
|
||||
// Create a real parent chat so the FK constraint is
|
||||
// satisfied.
|
||||
parentChat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "parent-chat",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{
|
||||
fantasy.TextContent{Text: "parent"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
ParentChatID: uuid.NullUUID{UUID: parentChat.ID, Valid: true},
|
||||
Title: "header-injection-parent",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{
|
||||
fantasy.TextContent{Text: "hello"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Wait for the chat to be processed and headers to be
|
||||
// captured.
|
||||
select {
|
||||
case <-headersCaptured:
|
||||
case <-ctx.Done():
|
||||
require.FailNow(t, "timed out waiting for SetExtraHeaders")
|
||||
}
|
||||
|
||||
require.Equal(t,
|
||||
chat.ID.String(),
|
||||
capturedHeaders.Get(workspacesdk.CoderChatIDHeader),
|
||||
)
|
||||
|
||||
ancestorJSON := capturedHeaders.Get(workspacesdk.CoderAncestorChatIDsHeader)
|
||||
var ancestorIDs []string
|
||||
err = json.Unmarshal([]byte(ancestorJSON), &ancestorIDs)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{parentChat.ID.String()}, ancestorIDs)
|
||||
})
|
||||
|
||||
t.Run("WithoutParentChat", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, model := seedChatDependencies(ctx, t, db)
|
||||
|
||||
org, err := db.GetDefaultOrganization(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
workspaceID, expectedAgentID := seedWorkspaceAgent(t, db, ps, user.ID, org.ID)
|
||||
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("done")...,
|
||||
)
|
||||
})
|
||||
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockConn := agentconnmock.NewMockAgentConn(ctrl)
|
||||
|
||||
var capturedHeaders http.Header
|
||||
headersCaptured := make(chan struct{})
|
||||
|
||||
mockConn.EXPECT().SetExtraHeaders(gomock.Any()).Do(func(h http.Header) {
|
||||
capturedHeaders = h
|
||||
close(headersCaptured)
|
||||
})
|
||||
mockConn.EXPECT().LS(gomock.Any(), gomock.Any(), gomock.Any()).Return(
|
||||
workspacesdk.LSResponse{}, xerrors.New("not found"),
|
||||
).AnyTimes()
|
||||
mockConn.EXPECT().Close().Return(nil).AnyTimes()
|
||||
|
||||
agentConnFn := func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
|
||||
require.Equal(t, expectedAgentID, agentID)
|
||||
return mockConn, func() {}, nil
|
||||
}
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := chatd.New(chatd.Config{
|
||||
Logger: logger,
|
||||
Database: db,
|
||||
ReplicaID: uuid.New(),
|
||||
Pubsub: ps,
|
||||
AgentConn: agentConnFn,
|
||||
PendingChatAcquireInterval: 10 * time.Millisecond,
|
||||
InFlightChatStaleAfter: testutil.WaitSuperLong,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, server.Close())
|
||||
})
|
||||
|
||||
// Create a chat without a parent — the ancestor header
|
||||
// should contain an empty JSON array.
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
Title: "header-injection-no-parent",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []fantasy.Content{
|
||||
fantasy.TextContent{Text: "hello"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
select {
|
||||
case <-headersCaptured:
|
||||
case <-ctx.Done():
|
||||
require.FailNow(t, "timed out waiting for SetExtraHeaders")
|
||||
}
|
||||
|
||||
require.Equal(t,
|
||||
chat.ID.String(),
|
||||
capturedHeaders.Get(workspacesdk.CoderChatIDHeader),
|
||||
)
|
||||
|
||||
// When there is no parent, the code declares
|
||||
// var ancestorIDs []string and never appends to it,
|
||||
// so json.Marshal produces "null".
|
||||
ancestorJSON := capturedHeaders.Get(workspacesdk.CoderAncestorChatIDsHeader)
|
||||
var ancestorIDs []string
|
||||
err = json.Unmarshal([]byte(ancestorJSON), &ancestorIDs)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, ancestorIDs)
|
||||
})
|
||||
}
|
||||
|
||||
+159
@@ -12,6 +12,7 @@ import (
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
@@ -35,6 +36,8 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/rbac/policy"
|
||||
"github.com/coder/coder/v2/coderd/tracing"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/wsjson"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -413,6 +416,162 @@ func (api *API) getChat(rw http.ResponseWriter, r *http.Request) {
|
||||
})
|
||||
}
|
||||
|
||||
// @Summary Watch git changes for a chat.
|
||||
// @ID watch-chat-git
|
||||
// @Security CoderSessionToken
|
||||
// @Tags Chats
|
||||
// @Param chat path string true "Chat ID" format(uuid)
|
||||
// @Success 101
|
||||
// @Router /chats/{chat}/git/watch [get]
|
||||
//
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
//
|
||||
//nolint:revive // HTTP handler writes to ResponseWriter.
|
||||
func (api *API) watchChatGit(rw http.ResponseWriter, r *http.Request) {
|
||||
var (
|
||||
ctx = r.Context()
|
||||
chat = httpmw.ChatParam(r)
|
||||
logger = api.Logger.Named("chat_git_watcher").With(slog.F("chat_id", chat.ID))
|
||||
)
|
||||
|
||||
if !chat.WorkspaceID.Valid {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Chat has no workspace to watch.",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
agents, err := api.Database.GetWorkspaceAgentsInLatestBuildByWorkspaceID(ctx, chat.WorkspaceID.UUID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error fetching workspace agents.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
if len(agents) == 0 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Chat workspace has no agents.",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
apiAgent, err := db2sdk.WorkspaceAgent(
|
||||
api.DERPMap(),
|
||||
*api.TailnetCoordinator.Load(),
|
||||
agents[0],
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
api.AgentInactiveDisconnectTimeout,
|
||||
api.DeploymentValues.AgentFallbackTroubleshootingURL.String(),
|
||||
)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error reading workspace agent.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
if apiAgent.Status != codersdk.WorkspaceAgentConnected {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: fmt.Sprintf("Agent state is %q, it must be in the %q state.", apiAgent.Status, codersdk.WorkspaceAgentConnected),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
dialCtx, dialCancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer dialCancel()
|
||||
|
||||
agentConn, release, err := api.agentProvider.AgentConn(dialCtx, agents[0].ID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error dialing workspace agent.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
defer release()
|
||||
|
||||
agentStream, err := agentConn.WatchGit(ctx, logger, chat.ID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error watching agent's git state.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
defer agentStream.Close(websocket.StatusGoingAway)
|
||||
|
||||
clientConn, err := websocket.Accept(rw, r, &websocket.AcceptOptions{
|
||||
CompressionMode: websocket.CompressionNoContextTakeover,
|
||||
})
|
||||
if err != nil {
|
||||
logger.Error(ctx, "failed to accept websocket", slog.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
clientStream := wsjson.NewStream[
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
](clientConn, websocket.MessageText, websocket.MessageText, logger)
|
||||
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
defer cancel()
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, clientConn)
|
||||
|
||||
// Proxy agent → client.
|
||||
agentCh := agentStream.Chan()
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for {
|
||||
select {
|
||||
case <-api.ctx.Done():
|
||||
return
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case msg, ok := <-agentCh:
|
||||
if !ok {
|
||||
cancel()
|
||||
return
|
||||
}
|
||||
if err := clientStream.Send(msg); err != nil {
|
||||
logger.Debug(ctx, "failed to forward agent message to client", slog.Error(err))
|
||||
cancel()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// Proxy client → agent.
|
||||
clientCh := clientStream.Chan()
|
||||
proxyLoop:
|
||||
for {
|
||||
select {
|
||||
case <-api.ctx.Done():
|
||||
break proxyLoop
|
||||
case <-ctx.Done():
|
||||
break proxyLoop
|
||||
case msg, ok := <-clientCh:
|
||||
if !ok {
|
||||
break proxyLoop
|
||||
}
|
||||
if err := agentStream.Send(msg); err != nil {
|
||||
logger.Debug(ctx, "failed to forward client message to agent", slog.Error(err))
|
||||
break proxyLoop
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
cancel()
|
||||
wg.Wait()
|
||||
_ = clientStream.Close(websocket.StatusGoingAway)
|
||||
}
|
||||
|
||||
// @Summary Archive a chat
|
||||
// @ID archive-chat
|
||||
// @Tags Chats
|
||||
|
||||
@@ -1132,6 +1132,7 @@ func New(options *Options) *API {
|
||||
r.Route("/{chat}", func(r chi.Router) {
|
||||
r.Use(httpmw.ExtractChatParam(options.Database))
|
||||
r.Get("/", api.getChat)
|
||||
r.Get("/git/watch", api.watchChatGit)
|
||||
r.Post("/archive", api.archiveChat)
|
||||
r.Post("/unarchive", api.unarchiveChat)
|
||||
r.Post("/messages", api.postChatMessages)
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"net/http/httputil"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
@@ -68,6 +69,526 @@ func (c *channelCloser) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestWatchChatGit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("ChatWithNoWorkspaceReturns400", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This test ensures that a chat with no workspace ID
|
||||
// returns a 400 error.
|
||||
|
||||
var (
|
||||
ctx = testutil.Context(t, testutil.WaitShort)
|
||||
logger = slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug).Named("coderd")
|
||||
|
||||
mCtrl = gomock.NewController(t)
|
||||
mDB = dbmock.NewMockStore(mCtrl)
|
||||
|
||||
chatID = uuid.New()
|
||||
|
||||
r = chi.NewMux()
|
||||
|
||||
api = API{
|
||||
ctx: ctx,
|
||||
Options: &Options{
|
||||
AgentInactiveDisconnectTimeout: testutil.WaitShort,
|
||||
Database: mDB,
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
// Setup: Return a chat with no workspace ID.
|
||||
mDB.EXPECT().GetChatByID(gomock.Any(), chatID).Return(database.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: uuid.New(),
|
||||
WorkspaceID: uuid.NullUUID{Valid: false},
|
||||
}, nil)
|
||||
|
||||
// And: We mount the HTTP handler.
|
||||
r.With(httpmw.ExtractChatParam(mDB)).
|
||||
Get("/chats/{chat}/git/watch", api.watchChatGit)
|
||||
|
||||
// Given: We create the HTTP server.
|
||||
srv := httptest.NewServer(r)
|
||||
defer srv.Close()
|
||||
|
||||
// When: We make a request.
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
||||
fmt.Sprintf("%s/chats/%s/git/watch", srv.URL, chatID), nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Then: We expect a 400 response.
|
||||
require.Equal(t, http.StatusBadRequest, resp.StatusCode)
|
||||
})
|
||||
|
||||
t.Run("UnauthorizedUsersCannotWatch", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This test ensures that if the chat middleware returns
|
||||
// an error (e.g. unauthorized), the handler is not
|
||||
// reached.
|
||||
|
||||
var (
|
||||
ctx = testutil.Context(t, testutil.WaitShort)
|
||||
logger = slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug).Named("coderd")
|
||||
|
||||
mCtrl = gomock.NewController(t)
|
||||
mDB = dbmock.NewMockStore(mCtrl)
|
||||
|
||||
chatID = uuid.New()
|
||||
|
||||
r = chi.NewMux()
|
||||
|
||||
api = API{
|
||||
ctx: ctx,
|
||||
Options: &Options{
|
||||
AgentInactiveDisconnectTimeout: testutil.WaitShort,
|
||||
Database: mDB,
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
// Setup: Return an error from the DB to simulate
|
||||
// unauthorized access.
|
||||
mDB.EXPECT().GetChatByID(gomock.Any(), chatID).Return(
|
||||
database.Chat{}, sql.ErrNoRows,
|
||||
)
|
||||
|
||||
// And: We mount the HTTP handler.
|
||||
r.With(httpmw.ExtractChatParam(mDB)).
|
||||
Get("/chats/{chat}/git/watch", api.watchChatGit)
|
||||
|
||||
// Given: We create the HTTP server.
|
||||
srv := httptest.NewServer(r)
|
||||
defer srv.Close()
|
||||
|
||||
// When: We make a request.
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
||||
fmt.Sprintf("%s/chats/%s/git/watch", srv.URL, chatID), nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Then: We expect a 404 (not found) since sql.ErrNoRows
|
||||
// is treated as a 404 by httpapi.Is404Error.
|
||||
require.Equal(t, http.StatusNotFound, resp.StatusCode)
|
||||
})
|
||||
|
||||
t.Run("DisconnectedAgentRejected", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This test ensures that a chat whose workspace agent is
|
||||
// not connected returns a 400 error.
|
||||
|
||||
var (
|
||||
ctx = testutil.Context(t, testutil.WaitShort)
|
||||
logger = slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug).Named("coderd")
|
||||
|
||||
mCtrl = gomock.NewController(t)
|
||||
mDB = dbmock.NewMockStore(mCtrl)
|
||||
mCoordinator = tailnettest.NewMockCoordinator(mCtrl)
|
||||
|
||||
chatID = uuid.New()
|
||||
workspaceID = uuid.New()
|
||||
agentID = uuid.New()
|
||||
resourceID = uuid.New()
|
||||
|
||||
r = chi.NewMux()
|
||||
|
||||
api = API{
|
||||
ctx: ctx,
|
||||
Options: &Options{
|
||||
AgentInactiveDisconnectTimeout: testutil.WaitShort,
|
||||
Database: mDB,
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
TailnetCoordinator: tailnettest.NewFakeCoordinator(),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
var tailnetCoordinator tailnet.Coordinator = mCoordinator
|
||||
api.TailnetCoordinator.Store(&tailnetCoordinator)
|
||||
|
||||
// Setup: Return a chat with a valid workspace ID.
|
||||
mDB.EXPECT().GetChatByID(gomock.Any(), chatID).Return(database.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: uuid.New(),
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
}, nil)
|
||||
|
||||
// And: Return an agent that is disconnected (no
|
||||
// FirstConnectedAt).
|
||||
mDB.EXPECT().GetWorkspaceAgentsInLatestBuildByWorkspaceID(gomock.Any(), workspaceID).
|
||||
Return([]database.WorkspaceAgent{{
|
||||
ID: agentID,
|
||||
ResourceID: resourceID,
|
||||
LifecycleState: database.WorkspaceAgentLifecycleStateCreated,
|
||||
}}, nil)
|
||||
|
||||
// And: Allow db2sdk.WorkspaceAgent to complete.
|
||||
mCoordinator.EXPECT().Node(gomock.Any()).Return(nil)
|
||||
|
||||
// And: We mount the HTTP handler.
|
||||
r.With(httpmw.ExtractChatParam(mDB)).
|
||||
Get("/chats/{chat}/git/watch", api.watchChatGit)
|
||||
|
||||
// Given: We create the HTTP server.
|
||||
srv := httptest.NewServer(r)
|
||||
defer srv.Close()
|
||||
|
||||
// When: We make a request.
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
||||
fmt.Sprintf("%s/chats/%s/git/watch", srv.URL, chatID), nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Then: We expect a 400 response since the agent is
|
||||
// not connected.
|
||||
require.Equal(t, http.StatusBadRequest, resp.StatusCode)
|
||||
})
|
||||
|
||||
t.Run("BidirectionalProxyWorks", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This test ensures that messages flow bidirectionally
|
||||
// between the client websocket and the agent websocket
|
||||
// through the coderd proxy.
|
||||
|
||||
var (
|
||||
ctx = testutil.Context(t, testutil.WaitLong)
|
||||
logger = slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug).Named("coderd")
|
||||
|
||||
mCtrl = gomock.NewController(t)
|
||||
mDB = dbmock.NewMockStore(mCtrl)
|
||||
mCoordinator = tailnettest.NewMockCoordinator(mCtrl)
|
||||
mAgentConn = agentconnmock.NewMockAgentConn(mCtrl)
|
||||
|
||||
chatID = uuid.New()
|
||||
workspaceID = uuid.New()
|
||||
agentID = uuid.New()
|
||||
resourceID = uuid.New()
|
||||
|
||||
r = chi.NewMux()
|
||||
|
||||
fAgentProvider = fakeAgentProvider{
|
||||
agentConn: func(ctx context.Context, aID uuid.UUID) (_ workspacesdk.AgentConn, release func(), _ error) {
|
||||
return mAgentConn, func() {}, nil
|
||||
},
|
||||
}
|
||||
|
||||
api = API{
|
||||
ctx: ctx,
|
||||
Options: &Options{
|
||||
AgentInactiveDisconnectTimeout: testutil.WaitShort,
|
||||
Database: mDB,
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
TailnetCoordinator: tailnettest.NewFakeCoordinator(),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
var tailnetCoordinator tailnet.Coordinator = mCoordinator
|
||||
api.TailnetCoordinator.Store(&tailnetCoordinator)
|
||||
api.agentProvider = fAgentProvider
|
||||
|
||||
// Setup: Create a fake agent-side websocket server that
|
||||
// we can interact with.
|
||||
agentDone := make(chan struct{})
|
||||
closeAgentDone := sync.OnceFunc(func() { close(agentDone) })
|
||||
t.Cleanup(closeAgentDone)
|
||||
agentStreamReady := make(chan *wsjson.Stream[
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
], 1)
|
||||
agentSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ws, err := websocket.Accept(w, r, nil)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
// Create stream typed from the agent's perspective:
|
||||
// reads client messages, writes server messages.
|
||||
s := wsjson.NewStream[
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
](ws, websocket.MessageText, websocket.MessageText, logger)
|
||||
agentStreamReady <- s
|
||||
// Keep the handler alive until test signals done.
|
||||
<-agentDone
|
||||
}))
|
||||
defer agentSrv.Close()
|
||||
|
||||
// And: Return a chat with a valid workspace ID.
|
||||
mDB.EXPECT().GetChatByID(gomock.Any(), chatID).Return(database.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: uuid.New(),
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
}, nil)
|
||||
|
||||
// And: Return a connected agent.
|
||||
mDB.EXPECT().GetWorkspaceAgentsInLatestBuildByWorkspaceID(gomock.Any(), workspaceID).
|
||||
Return([]database.WorkspaceAgent{{
|
||||
ID: agentID,
|
||||
ResourceID: resourceID,
|
||||
LifecycleState: database.WorkspaceAgentLifecycleStateReady,
|
||||
FirstConnectedAt: sql.NullTime{Valid: true, Time: dbtime.Now()},
|
||||
LastConnectedAt: sql.NullTime{Valid: true, Time: dbtime.Now()},
|
||||
}}, nil)
|
||||
|
||||
// And: Allow db2sdk.WorkspaceAgent to complete.
|
||||
mCoordinator.EXPECT().Node(gomock.Any()).Return(nil)
|
||||
|
||||
// And: WatchGit dials our fake agent server and returns
|
||||
// the stream.
|
||||
mAgentConn.EXPECT().WatchGit(gomock.Any(), gomock.Any(), chatID).
|
||||
DoAndReturn(func(ctx context.Context, _ slog.Logger, _ uuid.UUID) (*wsjson.Stream[codersdk.WorkspaceAgentGitServerMessage, codersdk.WorkspaceAgentGitClientMessage], error) {
|
||||
agentURL := strings.Replace(agentSrv.URL, "http://", "ws://", 1)
|
||||
conn, resp, err := websocket.Dial(ctx, agentURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp != nil && resp.Body != nil {
|
||||
defer resp.Body.Close()
|
||||
}
|
||||
// From coderd's perspective: reads server messages
|
||||
// from agent, writes client messages to agent.
|
||||
s := wsjson.NewStream[
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
](conn, websocket.MessageText, websocket.MessageText, logger)
|
||||
return s, nil
|
||||
})
|
||||
// And: We mount the HTTP handler.
|
||||
r.With(httpmw.ExtractChatParam(mDB)).
|
||||
Get("/chats/{chat}/git/watch", api.watchChatGit)
|
||||
|
||||
// Given: We create the HTTP server.
|
||||
coderdSrv := httptest.NewServer(r)
|
||||
defer coderdSrv.Close()
|
||||
|
||||
// And: Dial the WebSocket as a client.
|
||||
wsURL := strings.Replace(coderdSrv.URL, "http://", "ws://", 1)
|
||||
clientConn, resp, err := websocket.Dial(ctx, fmt.Sprintf("%s/chats/%s/git/watch", wsURL, chatID), nil)
|
||||
require.NoError(t, err)
|
||||
if resp.Body != nil {
|
||||
defer resp.Body.Close()
|
||||
}
|
||||
|
||||
// And: Create a client stream.
|
||||
clientStream := wsjson.NewStream[
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
](clientConn, websocket.MessageText, websocket.MessageText, logger)
|
||||
clientCh := clientStream.Chan()
|
||||
|
||||
// And: Wait for the agent stream to be ready.
|
||||
agentStream := testutil.RequireReceive(ctx, t, agentStreamReady)
|
||||
|
||||
// Test agent → client: Send a server message from the
|
||||
// agent and verify the client receives it.
|
||||
err = agentStream.Send(codersdk.WorkspaceAgentGitServerMessage{
|
||||
Type: codersdk.WorkspaceAgentGitServerMessageTypeChanges,
|
||||
Message: "test-changes",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
serverMsg := testutil.RequireReceive(ctx, t, clientCh)
|
||||
require.Equal(t, codersdk.WorkspaceAgentGitServerMessageTypeChanges, serverMsg.Type)
|
||||
require.Equal(t, "test-changes", serverMsg.Message)
|
||||
|
||||
// Test client → agent: Send a client message and verify
|
||||
// the agent receives it.
|
||||
agentCh := agentStream.Chan()
|
||||
err = clientStream.Send(codersdk.WorkspaceAgentGitClientMessage{
|
||||
Type: codersdk.WorkspaceAgentGitClientMessageTypeRefresh,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
clientMsg := testutil.RequireReceive(ctx, t, agentCh)
|
||||
require.Equal(t, codersdk.WorkspaceAgentGitClientMessageTypeRefresh, clientMsg.Type)
|
||||
|
||||
// Cleanup: Close the client connection to unwind the
|
||||
// proxy chain before closing the servers.
|
||||
_ = clientStream.Close(websocket.StatusNormalClosure)
|
||||
closeAgentDone()
|
||||
coderdSrv.Close()
|
||||
agentSrv.Close()
|
||||
})
|
||||
|
||||
t.Run("ClientDisconnectTearsDown", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This test ensures that closing the client websocket
|
||||
// causes the agent stream to be closed.
|
||||
|
||||
var (
|
||||
ctx = testutil.Context(t, testutil.WaitLong)
|
||||
logger = slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug).Named("coderd")
|
||||
|
||||
mCtrl = gomock.NewController(t)
|
||||
mDB = dbmock.NewMockStore(mCtrl)
|
||||
mCoordinator = tailnettest.NewMockCoordinator(mCtrl)
|
||||
mAgentConn = agentconnmock.NewMockAgentConn(mCtrl)
|
||||
|
||||
chatID = uuid.New()
|
||||
workspaceID = uuid.New()
|
||||
agentID = uuid.New()
|
||||
resourceID = uuid.New()
|
||||
|
||||
r = chi.NewMux()
|
||||
|
||||
fAgentProvider = fakeAgentProvider{
|
||||
agentConn: func(ctx context.Context, aID uuid.UUID) (_ workspacesdk.AgentConn, release func(), _ error) {
|
||||
return mAgentConn, func() {}, nil
|
||||
},
|
||||
}
|
||||
|
||||
api = API{
|
||||
ctx: ctx,
|
||||
Options: &Options{
|
||||
AgentInactiveDisconnectTimeout: testutil.WaitShort,
|
||||
Database: mDB,
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
TailnetCoordinator: tailnettest.NewFakeCoordinator(),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
var tailnetCoordinator tailnet.Coordinator = mCoordinator
|
||||
api.TailnetCoordinator.Store(&tailnetCoordinator)
|
||||
api.agentProvider = fAgentProvider
|
||||
|
||||
// Setup: Create a fake agent-side websocket server.
|
||||
agentDone := make(chan struct{})
|
||||
closeAgentDone := sync.OnceFunc(func() { close(agentDone) })
|
||||
t.Cleanup(closeAgentDone)
|
||||
agentStreamReady := make(chan *wsjson.Stream[
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
], 1)
|
||||
agentSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ws, err := websocket.Accept(w, r, nil)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
s := wsjson.NewStream[
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
](ws, websocket.MessageText, websocket.MessageText, logger)
|
||||
agentStreamReady <- s
|
||||
// Keep the handler alive until test signals done.
|
||||
<-agentDone
|
||||
}))
|
||||
defer agentSrv.Close()
|
||||
|
||||
// And: Return a chat with a valid workspace ID.
|
||||
mDB.EXPECT().GetChatByID(gomock.Any(), chatID).Return(database.Chat{
|
||||
ID: chatID,
|
||||
OwnerID: uuid.New(),
|
||||
WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true},
|
||||
}, nil)
|
||||
|
||||
// And: Return a connected agent.
|
||||
mDB.EXPECT().GetWorkspaceAgentsInLatestBuildByWorkspaceID(gomock.Any(), workspaceID).
|
||||
Return([]database.WorkspaceAgent{{
|
||||
ID: agentID,
|
||||
ResourceID: resourceID,
|
||||
LifecycleState: database.WorkspaceAgentLifecycleStateReady,
|
||||
FirstConnectedAt: sql.NullTime{Valid: true, Time: dbtime.Now()},
|
||||
LastConnectedAt: sql.NullTime{Valid: true, Time: dbtime.Now()},
|
||||
}}, nil)
|
||||
|
||||
// And: Allow db2sdk.WorkspaceAgent to complete.
|
||||
mCoordinator.EXPECT().Node(gomock.Any()).Return(nil)
|
||||
|
||||
// And: WatchGit dials our fake agent server.
|
||||
mAgentConn.EXPECT().WatchGit(gomock.Any(), gomock.Any(), chatID).
|
||||
DoAndReturn(func(ctx context.Context, _ slog.Logger, _ uuid.UUID) (*wsjson.Stream[codersdk.WorkspaceAgentGitServerMessage, codersdk.WorkspaceAgentGitClientMessage], error) {
|
||||
agentURL := strings.Replace(agentSrv.URL, "http://", "ws://", 1)
|
||||
conn, resp, err := websocket.Dial(ctx, agentURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp != nil && resp.Body != nil {
|
||||
defer resp.Body.Close()
|
||||
}
|
||||
s := wsjson.NewStream[
|
||||
codersdk.WorkspaceAgentGitServerMessage,
|
||||
codersdk.WorkspaceAgentGitClientMessage,
|
||||
](conn, websocket.MessageText, websocket.MessageText, logger)
|
||||
return s, nil
|
||||
})
|
||||
// And: We mount the HTTP handler.
|
||||
r.With(httpmw.ExtractChatParam(mDB)).
|
||||
Get("/chats/{chat}/git/watch", api.watchChatGit)
|
||||
|
||||
// Given: We create the HTTP server.
|
||||
coderdSrv := httptest.NewServer(r)
|
||||
defer coderdSrv.Close()
|
||||
|
||||
// And: Dial the WebSocket as a client.
|
||||
wsURL := strings.Replace(coderdSrv.URL, "http://", "ws://", 1)
|
||||
clientConn, resp, err := websocket.Dial(ctx, fmt.Sprintf("%s/chats/%s/git/watch", wsURL, chatID), nil)
|
||||
require.NoError(t, err)
|
||||
if resp.Body != nil {
|
||||
defer resp.Body.Close()
|
||||
}
|
||||
|
||||
// And: Wait for the agent stream to be ready.
|
||||
agentStream := testutil.RequireReceive(ctx, t, agentStreamReady)
|
||||
agentCh := agentStream.Chan()
|
||||
|
||||
// And: Verify the proxy is working first by sending a
|
||||
// message from agent to client.
|
||||
err = agentStream.Send(codersdk.WorkspaceAgentGitServerMessage{
|
||||
Type: codersdk.WorkspaceAgentGitServerMessageTypeChanges,
|
||||
Message: "hello",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
clientDecoder := wsjson.NewDecoder[codersdk.WorkspaceAgentGitServerMessage](clientConn, websocket.MessageText, logger)
|
||||
decodeCh := clientDecoder.Chan()
|
||||
serverMsg := testutil.RequireReceive(ctx, t, decodeCh)
|
||||
require.Equal(t, "hello", serverMsg.Message)
|
||||
|
||||
// When: We close the client WebSocket.
|
||||
clientConn.Close(websocket.StatusNormalClosure, "test closing connection")
|
||||
|
||||
// Then: We expect agentCh to be closed, indicating
|
||||
// teardown propagated to the agent side.
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for agent channel to close")
|
||||
|
||||
case _, ok := <-agentCh:
|
||||
require.False(t, ok, "agent channel is expected to be closed")
|
||||
}
|
||||
|
||||
// Cleanup: Close the servers in the correct order.
|
||||
closeAgentDone()
|
||||
coderdSrv.Close()
|
||||
agentSrv.Close()
|
||||
})
|
||||
}
|
||||
|
||||
func TestWatchAgentContainers(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user