fix: resolve bugs in chatd streaming system (#22720)

Split from #22693 per review feedback.

Fixes multiple bugs in coderd/chatd and sub-packages including race
conditions, transaction safety, stream buffer bounds, retry limits, and
enterprise relay improvements.

See commit message for full list.
This commit is contained in:
Kyle Carberry
2026-03-06 21:02:25 +00:00
committed by GitHub
parent 2cd871e88f
commit eecb7d0b66
7 changed files with 529 additions and 635 deletions
+20 -347
View File
@@ -6,7 +6,6 @@ import (
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"sync"
"sync/atomic"
@@ -17,8 +16,6 @@ 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"
@@ -32,8 +29,6 @@ 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"
@@ -78,10 +73,11 @@ func TestInterruptChatBroadcastsStatusAcrossInstances(t *testing.T) {
require.Eventually(t, func() bool {
select {
case event := <-events:
if event.Type != codersdk.ChatStreamEventTypeStatus || event.Status == nil {
return false
if event.Type == codersdk.ChatStreamEventTypeStatus && event.Status != nil {
return event.Status.Status == codersdk.ChatStatusWaiting
}
return event.Status.Status == codersdk.ChatStatusWaiting
t.Logf("skipping unexpected event: type=%s", event.Type)
return false
default:
return false
}
@@ -870,15 +866,15 @@ func TestSubscribeNoPubsubNoDuplicateMessageParts(t *testing.T) {
// events — the snapshot already contained everything. Before
// the fix, localSnapshot was replayed into the channel,
// causing duplicates.
select {
case event, ok := <-events:
if ok {
t.Fatalf("unexpected event from channel (would be a duplicate): type=%s", event.Type)
require.Never(t, func() bool {
select {
case <-events:
return true
default:
return false
}
// Channel closed without events is fine.
case <-time.After(200 * time.Millisecond):
// No events — correct behavior.
}
}, 200*time.Millisecond, testutil.IntervalFast,
"expected no duplicate events after snapshot")
}
func TestSubscribeAfterMessageID(t *testing.T) {
@@ -1533,13 +1529,16 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) {
})
require.NoError(t, err)
// Wait for a web push notification to be dispatched. The dispatch
// happens asynchronously after the DB status is updated, so we need
// to poll rather than assert immediately.
testutil.Eventually(ctx, t, func(_ context.Context) bool {
return mockPush.dispatchCount.Load() >= 1
// Wait for the chat to complete and return to waiting status.
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
fromDB, dbErr := db.GetChatByID(ctx, chat.ID)
if dbErr != nil {
return false
}
return fromDB.Status == database.ChatStatusWaiting && !fromDB.WorkerID.Valid && mockPush.dispatchCount.Load() == 1
}, testutil.IntervalFast)
// Verify a web push notification was dispatched exactly once.
require.Equal(t, int32(1), mockPush.dispatchCount.Load(),
"expected exactly one web push dispatch for a completed chat")
@@ -1558,75 +1557,6 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) {
"web push Data should contain the chat navigation URL")
}
func TestSuccessfulChatSendsWebPushWithTag(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
// Set up a mock OpenAI that returns a simple streaming response.
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
if !req.Stream {
return chattest.OpenAINonStreamingResponse("title")
}
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("done")...)
})
// Mock webpush dispatcher that captures calls.
mockPush := &mockWebpushDispatcher{}
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
server := chatd.New(chatd.Config{
Logger: logger,
Database: db,
ReplicaID: uuid.New(),
Pubsub: ps,
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
})
t.Cleanup(func() {
require.NoError(t, server.Close())
})
user, model := seedChatDependencies(ctx, t, db)
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OwnerID: user.ID,
Title: "push-tag-test",
ModelConfigID: model.ID,
InitialUserContent: []fantasy.Content{fantasy.TextContent{Text: "hello"}},
})
require.NoError(t, err)
// Wait for the web push notification to be dispatched.
// We poll dispatchCount rather than DB status because the
// push fires after the status update, creating a small race
// window.
testutil.Eventually(ctx, t, func(_ context.Context) bool {
return mockPush.dispatchCount.Load() >= 1
}, testutil.IntervalFast)
require.Equal(t, int32(1), mockPush.dispatchCount.Load(),
"expected exactly one web push dispatch for a completed chat")
// Verify the push notification tag is set to the chat ID for dedup.
mockPush.mu.Lock()
capturedMsg := mockPush.lastMessage
capturedUser := mockPush.lastUserID
mockPush.mu.Unlock()
require.Equal(t, chat.ID.String(), capturedMsg.Tag,
"push notification tag should equal the chat ID for deduplication")
require.Equal(t, user.ID, capturedUser,
"push notification should be dispatched to the chat owner")
require.Equal(t, "push-tag-test", capturedMsg.Title,
"push notification title should match the chat title")
require.Equal(t, "Agent has finished running.", capturedMsg.Body,
"push notification body should indicate the agent finished")
}
func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T) {
t.Parallel()
@@ -1733,260 +1663,3 @@ 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)
})
}