diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index 51874b7fd7..b79ccc12eb 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -492,6 +492,43 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C } } + var workspaceAwareness string + if opts.WorkspaceID.Valid { + workspaceAwareness = "This chat is attached to a workspace. You can use workspace tools like execute, read_file, write_file, etc." + } else { + workspaceAwareness = "There is no workspace associated with this chat yet. Create one using the create_workspace tool before using workspace tools like execute, read_file, write_file, etc." + } + workspaceAwarenessContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText(workspaceAwareness), + }) + if err != nil { + return xerrors.Errorf("marshal workspace awareness: %w", err) + } + _, err = tx.InsertChatMessage(ctx, database.InsertChatMessageParams{ + ChatID: insertedChat.ID, + CreatedBy: uuid.NullUUID{}, + ModelConfigID: uuid.NullUUID{ + UUID: opts.ModelConfigID, + Valid: true, + }, + Role: database.ChatMessageRoleSystem, + ContentVersion: chatprompt.CurrentContentVersion, + Content: workspaceAwarenessContent, + Visibility: database.ChatMessageVisibilityModel, + InputTokens: sql.NullInt64{}, + OutputTokens: sql.NullInt64{}, + TotalTokens: sql.NullInt64{}, + ReasoningTokens: sql.NullInt64{}, + CacheCreationTokens: sql.NullInt64{}, + CacheReadTokens: sql.NullInt64{}, + ContextLimit: sql.NullInt64{}, + Compressed: sql.NullBool{}, + TotalCostMicros: sql.NullInt64{}, + }) + if err != nil { + return xerrors.Errorf("insert workspace awareness message: %w", err) + } + userContent, err := chatprompt.MarshalParts(opts.InitialUserContent) if err != nil { return xerrors.Errorf("marshal initial user content: %w", err) diff --git a/coderd/chatd/chatd_test.go b/coderd/chatd/chatd_test.go index 7f5d71874d..f6b07047ec 100644 --- a/coderd/chatd/chatd_test.go +++ b/coderd/chatd/chatd_test.go @@ -591,6 +591,97 @@ func TestEditMessageUpdatesAndTruncatesAndClearsQueue(t *testing.T) { require.False(t, chatFromDB.WorkerID.Valid) } +func TestCreateChatInsertsWorkspaceAwarenessMessage(t *testing.T) { + t.Parallel() + + t.Run("WithWorkspace", func(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + server := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + + org := dbgen.Organization(t, db, database.Organization{}) + tv := dbgen.TemplateVersion(t, db, database.TemplateVersion{ + OrganizationID: org.ID, + CreatedBy: user.ID, + }) + tpl := dbgen.Template(t, db, database.Template{ + CreatedBy: user.ID, + OrganizationID: org.ID, + ActiveVersionID: tv.ID, + }) + workspace := dbgen.Workspace(t, db, database.WorkspaceTable{ + OwnerID: user.ID, + OrganizationID: org.ID, + TemplateID: tpl.ID, + }) + + chat, err := server.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + WorkspaceID: uuid.NullUUID{UUID: workspace.ID, Valid: true}, + Title: "test-with-workspace", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + messages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID) + require.NoError(t, err) + + var workspaceMsg *database.ChatMessage + for _, msg := range messages { + if msg.Role == database.ChatMessageRoleSystem { + content := string(msg.Content.RawMessage) + if strings.Contains(content, "attached to a workspace") { + workspaceMsg = &msg + break + } + } + } + require.NotNil(t, workspaceMsg, "workspace awareness system message should exist") + require.Equal(t, database.ChatMessageRoleSystem, workspaceMsg.Role) + require.Equal(t, database.ChatMessageVisibilityModel, workspaceMsg.Visibility) + }) + + t.Run("WithoutWorkspace", func(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + server := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, model := seedChatDependencies(ctx, t, db) + + chat, err := server.CreateChat(ctx, chatd.CreateOptions{ + OwnerID: user.ID, + Title: "test-without-workspace", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + messages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID) + require.NoError(t, err) + + var workspaceMsg *database.ChatMessage + for _, msg := range messages { + if msg.Role == database.ChatMessageRoleSystem { + content := string(msg.Content.RawMessage) + if strings.Contains(content, "no workspace associated") { + workspaceMsg = &msg + break + } + } + } + require.NotNil(t, workspaceMsg, "workspace awareness system message should exist") + require.Equal(t, database.ChatMessageRoleSystem, workspaceMsg.Role) + require.Equal(t, database.ChatMessageVisibilityModel, workspaceMsg.Visibility) + }) +} + func TestCreateChatRejectsWhenUsageLimitReached(t *testing.T) { t.Parallel()