refactor: add dbgen chat generators and migrate test boilerplate (#24497)

- Adds chat-related dbgen generators covering defaults, overrides, and message field mapping.
- Replaces raw single-row chat, message, provider, and model-config setup in tests with dbgen helpers.
- Simplifies chat seed helpers after moving fixture setup into dbgen.

> Generated with [Coder Agents](https://coder.com/agents).
This commit is contained in:
Cian Johnston
2026-05-01 13:29:33 +01:00
committed by GitHub
parent 6ee5fe983c
commit 2f855904be
25 changed files with 1586 additions and 2316 deletions
+161
View File
@@ -29,6 +29,7 @@ import (
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/rbac/policy"
"github.com/coder/coder/v2/coderd/rbac/rolestore"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/cryptorand"
"github.com/coder/coder/v2/provisionerd/proto"
@@ -75,6 +76,166 @@ func AuditLog(t testing.TB, db database.Store, seed database.AuditLog) database.
return log
}
func Chat(t testing.TB, db database.Store, seed database.Chat) database.Chat {
t.Helper()
var labels pqtype.NullRawMessage
if seed.Labels != nil {
raw, err := json.Marshal(seed.Labels)
require.NoError(t, err, "marshal chat labels")
labels = pqtype.NullRawMessage{RawMessage: raw, Valid: true}
}
chat, err := db.InsertChat(genCtx, database.InsertChatParams{
OrganizationID: takeFirst(seed.OrganizationID, uuid.New()),
OwnerID: takeFirst(seed.OwnerID, uuid.New()),
WorkspaceID: seed.WorkspaceID,
BuildID: seed.BuildID,
AgentID: seed.AgentID,
ParentChatID: seed.ParentChatID,
RootChatID: seed.RootChatID,
LastModelConfigID: takeFirst(seed.LastModelConfigID, uuid.New()),
Title: takeFirst(seed.Title, testutil.GetRandomName(t)),
Mode: seed.Mode,
PlanMode: seed.PlanMode,
Status: takeFirst(seed.Status, database.ChatStatusWaiting),
MCPServerIDs: seed.MCPServerIDs,
Labels: labels,
DynamicTools: seed.DynamicTools,
ClientType: takeFirst(seed.ClientType, database.ChatClientTypeUi),
})
require.NoError(t, err, "insert chat")
return chat
}
func ChatMessage(t testing.TB, db database.Store, seed database.ChatMessage) database.ChatMessage {
t.Helper()
content := "[]"
if seed.Content.Valid {
content = string(seed.Content.RawMessage)
}
msgs, err := db.InsertChatMessages(genCtx, database.InsertChatMessagesParams{
ChatID: seed.ChatID,
CreatedBy: []uuid.UUID{seed.CreatedBy.UUID},
ModelConfigID: []uuid.UUID{seed.ModelConfigID.UUID},
Role: []database.ChatMessageRole{takeFirst(seed.Role, database.ChatMessageRoleUser)},
Content: []string{content},
ContentVersion: []int16{takeFirst(seed.ContentVersion, chatprompt.CurrentContentVersion)},
Visibility: []database.ChatMessageVisibility{takeFirst(seed.Visibility, database.ChatMessageVisibilityBoth)},
InputTokens: []int64{seed.InputTokens.Int64},
OutputTokens: []int64{seed.OutputTokens.Int64},
TotalTokens: []int64{seed.TotalTokens.Int64},
ReasoningTokens: []int64{seed.ReasoningTokens.Int64},
CacheCreationTokens: []int64{seed.CacheCreationTokens.Int64},
CacheReadTokens: []int64{seed.CacheReadTokens.Int64},
ContextLimit: []int64{seed.ContextLimit.Int64},
Compressed: []bool{seed.Compressed},
TotalCostMicros: []int64{seed.TotalCostMicros.Int64},
RuntimeMs: []int64{seed.RuntimeMs.Int64},
ProviderResponseID: []string{seed.ProviderResponseID.String},
})
require.NoError(t, err, "insert chat message")
require.Len(t, msgs, 1)
return msgs[0]
}
const (
// Match the default OpenAI test model's effective context settings.
defaultChatModelContextLimit int64 = 128000
defaultChatModelCompressionThreshold int32 = 70
)
func ChatModelConfig(t testing.TB, db database.Store, seed database.ChatModelConfig, munge ...func(*database.InsertChatModelConfigParams)) database.ChatModelConfig {
t.Helper()
params := database.InsertChatModelConfigParams{
Provider: takeFirst(seed.Provider, "openai"),
Model: takeFirst(seed.Model, "gpt-4o-mini"),
DisplayName: takeFirst(seed.DisplayName, "Test Model"),
CreatedBy: seed.CreatedBy,
UpdatedBy: seed.UpdatedBy,
Enabled: takeFirst(seed.Enabled, true),
IsDefault: seed.IsDefault,
ContextLimit: takeFirst(seed.ContextLimit, defaultChatModelContextLimit),
CompressionThreshold: takeFirst(seed.CompressionThreshold, defaultChatModelCompressionThreshold),
Options: takeFirstSlice(seed.Options, json.RawMessage(`{}`)),
}
for _, fn := range munge {
fn(&params)
}
cfg, err := db.InsertChatModelConfig(genCtx, params)
require.NoError(t, err, "insert chat model config")
return cfg
}
func ChatProvider(t testing.TB, db database.Store, seed database.ChatProvider, munge ...func(*database.InsertChatProviderParams)) database.ChatProvider {
t.Helper()
params := database.InsertChatProviderParams{
Provider: takeFirst(seed.Provider, "openai"),
DisplayName: takeFirst(seed.DisplayName, seed.Provider, "openai"),
APIKey: takeFirst(seed.APIKey, "test-key"),
BaseUrl: seed.BaseUrl,
ApiKeyKeyID: seed.ApiKeyKeyID,
CreatedBy: seed.CreatedBy,
Enabled: takeFirst(seed.Enabled, true),
CentralApiKeyEnabled: takeFirst(seed.CentralApiKeyEnabled, true),
AllowUserApiKey: seed.AllowUserApiKey,
AllowCentralApiKeyFallback: seed.AllowCentralApiKeyFallback,
}
for _, fn := range munge {
fn(&params)
}
provider, err := db.InsertChatProvider(genCtx, params)
require.NoError(t, err, "insert chat provider")
return provider
}
func MCPServerConfig(t testing.TB, db database.Store, seed database.MCPServerConfig) database.MCPServerConfig {
t.Helper()
// CreatedBy and UpdatedBy are user FKs, so default fixtures create a user.
createdBy := seed.CreatedBy.UUID
if createdBy == uuid.Nil {
createdBy = User(t, db, database.User{}).ID
}
updatedBy := seed.UpdatedBy.UUID
if updatedBy == uuid.Nil {
updatedBy = createdBy
}
cfg, err := db.InsertMCPServerConfig(genCtx, database.InsertMCPServerConfigParams{
DisplayName: takeFirst(seed.DisplayName, "Test MCP Server"),
Slug: takeFirst(seed.Slug, testutil.GetRandomName(t)),
Description: seed.Description,
IconURL: seed.IconURL,
Transport: takeFirst(seed.Transport, "streamable_http"),
Url: takeFirst(seed.Url, "https://mcp.example.com"),
AuthType: takeFirst(seed.AuthType, "none"),
OAuth2ClientID: seed.OAuth2ClientID,
OAuth2ClientSecret: seed.OAuth2ClientSecret,
OAuth2ClientSecretKeyID: seed.OAuth2ClientSecretKeyID,
OAuth2AuthURL: seed.OAuth2AuthURL,
OAuth2TokenURL: seed.OAuth2TokenURL,
OAuth2Scopes: seed.OAuth2Scopes,
APIKeyHeader: seed.APIKeyHeader,
APIKeyValue: seed.APIKeyValue,
APIKeyValueKeyID: seed.APIKeyValueKeyID,
CustomHeaders: seed.CustomHeaders,
CustomHeadersKeyID: seed.CustomHeadersKeyID,
ToolAllowList: takeFirstSlice(seed.ToolAllowList, []string{}),
ToolDenyList: takeFirstSlice(seed.ToolDenyList, []string{}),
Availability: takeFirst(seed.Availability, "default_off"),
Enabled: takeFirst(seed.Enabled, true),
ModelIntent: seed.ModelIntent,
AllowInPlanMode: seed.AllowInPlanMode,
CreatedBy: createdBy,
UpdatedBy: updatedBy,
})
require.NoError(t, err, "insert MCP server config")
return cfg
}
func ConnectionLog(t testing.TB, db database.Store, seed database.UpsertConnectionLogParams) database.ConnectionLog {
arg := database.UpsertConnectionLogParams{
ID: takeFirst(seed.ID, uuid.New()),
+189
View File
@@ -2,14 +2,18 @@ package dbgen_test
import (
"context"
"database/sql"
"encoding/json"
"testing"
"github.com/google/uuid"
"github.com/sqlc-dev/pqtype"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
)
func TestGenerator(t *testing.T) {
@@ -252,6 +256,191 @@ func TestGenerator(t *testing.T) {
require.Len(t, actual, 1)
require.Equal(t, exp, actual[0])
})
t.Run("ChatProvider", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
// Defaults.
p := dbgen.ChatProvider(t, db, database.ChatProvider{})
require.NotEqual(t, uuid.Nil, p.ID)
require.Equal(t, "openai", p.Provider)
require.Equal(t, "openai", p.DisplayName)
require.True(t, p.Enabled)
require.True(t, p.CentralApiKeyEnabled)
require.Equal(t, "test-key", p.APIKey)
// Overrides.
p2 := dbgen.ChatProvider(t, db, database.ChatProvider{
Provider: "anthropic",
DisplayName: "Claude",
APIKey: "sk-custom",
})
require.Equal(t, "anthropic", p2.Provider)
require.Equal(t, "Claude", p2.DisplayName)
require.Equal(t, "sk-custom", p2.APIKey)
p3 := dbgen.ChatProvider(t, db, database.ChatProvider{
Provider: "openrouter",
}, func(params *database.InsertChatProviderParams) {
params.APIKey = ""
})
require.Empty(t, p3.APIKey)
})
t.Run("ChatModelConfig", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
_ = dbgen.ChatProvider(t, db, database.ChatProvider{})
// Defaults.
cfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{})
require.NotEqual(t, uuid.Nil, cfg.ID)
require.Equal(t, "openai", cfg.Provider)
require.Equal(t, "gpt-4o-mini", cfg.Model)
require.Equal(t, "Test Model", cfg.DisplayName)
require.True(t, cfg.Enabled)
require.Equal(t, int64(128000), cfg.ContextLimit)
require.Equal(t, int32(70), cfg.CompressionThreshold)
// Overrides.
_ = dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "anthropic"})
cfg2 := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
Provider: "anthropic",
Model: "claude-4",
ContextLimit: 200000,
})
require.Equal(t, "anthropic", cfg2.Provider)
require.Equal(t, "claude-4", cfg2.Model)
require.Equal(t, int64(200000), cfg2.ContextLimit)
})
t.Run("Chat", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
u := dbgen.User(t, db, database.User{})
o := dbgen.Organization(t, db, database.Organization{})
dbgen.OrganizationMember(t, db, database.OrganizationMember{
UserID: u.ID,
OrganizationID: o.ID,
})
p := dbgen.ChatProvider(t, db, database.ChatProvider{})
m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: p.Provider})
// Defaults.
chat := dbgen.Chat(t, db, database.Chat{
OwnerID: u.ID,
OrganizationID: o.ID,
LastModelConfigID: m.ID,
})
require.NotEqual(t, uuid.Nil, chat.ID)
require.Equal(t, database.ChatStatusWaiting, chat.Status)
require.Equal(t, database.ChatClientTypeUi, chat.ClientType)
require.NotEmpty(t, chat.Title)
// Overrides.
chat2 := dbgen.Chat(t, db, database.Chat{
OwnerID: u.ID,
OrganizationID: o.ID,
LastModelConfigID: m.ID,
Title: "custom-title",
Status: database.ChatStatusRunning,
})
require.Equal(t, "custom-title", chat2.Title)
require.Equal(t, database.ChatStatusRunning, chat2.Status)
})
t.Run("ChatMessage", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
u := dbgen.User(t, db, database.User{})
o := dbgen.Organization(t, db, database.Organization{})
dbgen.OrganizationMember(t, db, database.OrganizationMember{
UserID: u.ID,
OrganizationID: o.ID,
})
p := dbgen.ChatProvider(t, db, database.ChatProvider{})
m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: p.Provider})
chat := dbgen.Chat(t, db, database.Chat{
OwnerID: u.ID,
OrganizationID: o.ID,
LastModelConfigID: m.ID,
})
// Defaults.
msg := dbgen.ChatMessage(t, db, database.ChatMessage{
ChatID: chat.ID,
})
require.NotZero(t, msg.ID)
require.Equal(t, database.ChatMessageRoleUser, msg.Role)
require.Equal(t, database.ChatMessageVisibilityBoth, msg.Visibility)
require.Equal(t, chatprompt.CurrentContentVersion, msg.ContentVersion)
// Overrides.
rawContent := json.RawMessage(`[{"type":"text","text":"hello"}]`)
msg2 := dbgen.ChatMessage(t, db, database.ChatMessage{
ChatID: chat.ID,
Role: database.ChatMessageRoleAssistant,
Content: pqtype.NullRawMessage{
RawMessage: rawContent,
Valid: true,
},
InputTokens: sql.NullInt64{Int64: 11, Valid: true},
OutputTokens: sql.NullInt64{Int64: 22, Valid: true},
TotalTokens: sql.NullInt64{Int64: 33, Valid: true},
ReasoningTokens: sql.NullInt64{Int64: 44, Valid: true},
CacheCreationTokens: sql.NullInt64{Int64: 55, Valid: true},
CacheReadTokens: sql.NullInt64{Int64: 66, Valid: true},
ContextLimit: sql.NullInt64{Int64: 77, Valid: true},
Compressed: true,
TotalCostMicros: sql.NullInt64{Int64: 88, Valid: true},
ProviderResponseID: sql.NullString{String: "resp-123", Valid: true},
})
require.Equal(t, database.ChatMessageRoleAssistant, msg2.Role)
require.True(t, msg2.Content.Valid)
require.JSONEq(t, string(rawContent), string(msg2.Content.RawMessage))
require.Equal(t, sql.NullInt64{Int64: 11, Valid: true}, msg2.InputTokens)
require.Equal(t, sql.NullInt64{Int64: 22, Valid: true}, msg2.OutputTokens)
require.Equal(t, sql.NullInt64{Int64: 33, Valid: true}, msg2.TotalTokens)
require.Equal(t, sql.NullInt64{Int64: 44, Valid: true}, msg2.ReasoningTokens)
require.Equal(t, sql.NullInt64{Int64: 55, Valid: true}, msg2.CacheCreationTokens)
require.Equal(t, sql.NullInt64{Int64: 66, Valid: true}, msg2.CacheReadTokens)
require.Equal(t, sql.NullInt64{Int64: 77, Valid: true}, msg2.ContextLimit)
require.True(t, msg2.Compressed)
require.Equal(t, sql.NullInt64{Int64: 88, Valid: true}, msg2.TotalCostMicros)
require.Equal(t, sql.NullString{String: "resp-123", Valid: true}, msg2.ProviderResponseID)
})
t.Run("MCPServerConfig", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
// Defaults.
cfg := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{})
require.NotEqual(t, uuid.Nil, cfg.ID)
require.Equal(t, "streamable_http", cfg.Transport)
require.Equal(t, "none", cfg.AuthType)
require.Equal(t, "default_off", cfg.Availability)
require.True(t, cfg.Enabled)
require.Empty(t, cfg.ToolAllowList)
require.Empty(t, cfg.ToolDenyList)
require.NotEmpty(t, cfg.Slug)
require.NotEmpty(t, cfg.Url)
// Overrides.
cfg2 := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{
DisplayName: "Custom MCP",
Slug: "custom-mcp",
Url: "https://custom.example.com",
AuthType: "oauth2",
AllowInPlanMode: true,
})
require.Equal(t, "Custom MCP", cfg2.DisplayName)
require.Equal(t, "custom-mcp", cfg2.Slug)
require.Equal(t, "https://custom.example.com", cfg2.Url)
require.Equal(t, "oauth2", cfg2.AuthType)
require.True(t, cfg2.AllowInPlanMode)
})
}
func must[T any](value T, err error) T {
+36 -89
View File
@@ -1839,20 +1839,17 @@ func TestDeleteOldChatFiles(t *testing.T) {
// backdates updated_at to control the "archived since" window.
createChat := func(ctx context.Context, t *testing.T, db database.Store, rawDB *sql.DB, ownerID, orgID, modelConfigID uuid.UUID, archived bool, updatedAt time.Time) database.Chat {
t.Helper()
chat, err := db.InsertChat(ctx, database.InsertChatParams{
chat := dbgen.Chat(t, db, database.Chat{
OrganizationID: orgID,
OwnerID: ownerID,
LastModelConfigID: modelConfigID,
Title: "test-chat",
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
})
require.NoError(t, err)
if archived {
_, err = db.ArchiveChatByID(ctx, chat.ID)
_, err := db.ArchiveChatByID(ctx, chat.ID)
require.NoError(t, err)
}
_, err = rawDB.ExecContext(ctx, "UPDATE chats SET updated_at = $1 WHERE id = $2", updatedAt, chat.ID)
_, err := rawDB.ExecContext(ctx, "UPDATE chats SET updated_at = $1 WHERE id = $2", updatedAt, chat.ID)
require.NoError(t, err)
return chat
}
@@ -1863,25 +1860,20 @@ func TestDeleteOldChatFiles(t *testing.T) {
org database.Organization
modelConfig database.ChatModelConfig
}
setupChatDeps := func(ctx context.Context, t *testing.T, db database.Store) chatDeps {
setupChatDeps := func(t *testing.T, db database.Store) chatDeps {
t.Helper()
user := dbgen.User(t, db, database.User{})
org := dbgen.Organization(t, db, database.Organization{})
_ = dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: user.ID, OrganizationID: org.ID})
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
Provider: "openai",
DisplayName: "OpenAI",
Enabled: true,
CentralApiKeyEnabled: true,
_ = dbgen.ChatProvider(t, db, database.ChatProvider{
Provider: "openai",
DisplayName: "OpenAI",
})
require.NoError(t, err)
mc, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
mc := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
Provider: "openai",
Model: "test-model",
ContextLimit: 8192,
Options: json.RawMessage("{}"),
})
require.NoError(t, err)
return chatDeps{user: user, org: org, modelConfig: mc}
}
@@ -1898,7 +1890,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
// Disable retention.
err := db.UpsertChatRetentionDays(ctx, int32(0))
@@ -1929,7 +1921,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
err := db.UpsertChatRetentionDays(ctx, int32(30))
require.NoError(t, err)
@@ -1937,27 +1929,12 @@ func TestDeleteOldChatFiles(t *testing.T) {
// Old archived chat (31 days) — should be deleted.
oldChat := createChat(ctx, t, db, rawDB, deps.user.ID, deps.org.ID, deps.modelConfig.ID, true, now.Add(-31*24*time.Hour))
// Insert a message so we can verify CASCADE.
_, err = db.InsertChatMessages(ctx, database.InsertChatMessagesParams{
ChatID: oldChat.ID,
CreatedBy: []uuid.UUID{deps.user.ID},
ModelConfigID: []uuid.UUID{deps.modelConfig.ID},
Role: []database.ChatMessageRole{database.ChatMessageRoleUser},
Content: []string{`[{"type":"text","text":"hello"}]`},
ContentVersion: []int16{0},
Visibility: []database.ChatMessageVisibility{database.ChatMessageVisibilityBoth},
InputTokens: []int64{0},
OutputTokens: []int64{0},
TotalTokens: []int64{0},
ReasoningTokens: []int64{0},
CacheCreationTokens: []int64{0},
CacheReadTokens: []int64{0},
ContextLimit: []int64{0},
Compressed: []bool{false},
TotalCostMicros: []int64{0},
RuntimeMs: []int64{0},
ProviderResponseID: []string{""},
_ = dbgen.ChatMessage(t, db, database.ChatMessage{
ChatID: oldChat.ID,
CreatedBy: uuid.NullUUID{UUID: deps.user.ID, Valid: true},
ModelConfigID: uuid.NullUUID{UUID: deps.modelConfig.ID, Valid: true},
Role: database.ChatMessageRoleUser,
})
require.NoError(t, err)
// Recently archived chat (10 days) — should be retained.
recentChat := createChat(ctx, t, db, rawDB, deps.user.ID, deps.org.ID, deps.modelConfig.ID, true, now.Add(-10*24*time.Hour))
@@ -1998,7 +1975,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
err := db.UpsertChatRetentionDays(ctx, int32(30))
require.NoError(t, err)
@@ -2049,7 +2026,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
err := db.UpsertChatRetentionDays(ctx, int32(30))
require.NoError(t, err)
@@ -2126,7 +2103,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
// file purge should show only surviving files.
ctx := testutil.Context(t, testutil.WaitLong)
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
// Create a chat with three attached files.
fileA := createChatFile(ctx, t, db, rawDB, deps.user.ID, deps.org.ID, now)
@@ -2179,19 +2156,13 @@ func TestDeleteOldChatFiles(t *testing.T) {
// clean up links for both parent and child chats
// independently via FK cascade.
parentChat := createChat(ctx, t, db, rawDB, deps.user.ID, deps.org.ID, deps.modelConfig.ID, false, now)
childChat, err := db.InsertChat(ctx, database.InsertChatParams{
childChat := dbgen.Chat(t, db, database.Chat{
OrganizationID: deps.org.ID,
OwnerID: deps.user.ID,
LastModelConfigID: deps.modelConfig.ID,
RootChatID: uuid.NullUUID{UUID: parentChat.ID, Valid: true},
Title: "child-chat",
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
})
require.NoError(t, err)
// Set root_chat_id to link child to parent.
_, err = rawDB.ExecContext(ctx, "UPDATE chats SET root_chat_id = $1 WHERE id = $2", parentChat.ID, childChat.ID)
require.NoError(t, err)
// Attach different files to parent and child.
parentFileKeep := createChatFile(ctx, t, db, rawDB, deps.user.ID, deps.org.ID, now)
@@ -2243,7 +2214,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
run: func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
// Create 3 deletable orphaned files (all 31 days old).
for range 3 {
@@ -2272,7 +2243,7 @@ func TestDeleteOldChatFiles(t *testing.T) {
run: func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
db, _, rawDB := dbtestutil.NewDBWithSQLDB(t, dbtestutil.WithDumpOnFailure())
deps := setupChatDeps(ctx, t, db)
deps := setupChatDeps(t, db)
// Create 3 deletable old archived chats.
for range 3 {
@@ -2307,25 +2278,20 @@ func TestDeleteOldChatFiles(t *testing.T) {
// helpers for TestAutoArchiveInactiveChats. Kept scoped to the
// test so they don't leak into the package surface area.
func archiveTestDeps(ctx context.Context, t *testing.T, db database.Store) chatAutoArchiveDeps {
func archiveTestDeps(t *testing.T, db database.Store) chatAutoArchiveDeps {
t.Helper()
user := dbgen.User(t, db, database.User{})
org := dbgen.Organization(t, db, database.Organization{})
_ = dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: user.ID, OrganizationID: org.ID})
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
Provider: "openai",
DisplayName: "OpenAI",
Enabled: true,
CentralApiKeyEnabled: true,
_ = dbgen.ChatProvider(t, db, database.ChatProvider{
Provider: "openai",
DisplayName: "OpenAI",
})
require.NoError(t, err)
mc, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
mc := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
Provider: "openai",
Model: "test-model",
ContextLimit: 8192,
Options: json.RawMessage("{}"),
})
require.NoError(t, err)
return chatAutoArchiveDeps{user: user, org: org, modelConfig: mc}
}
@@ -2361,7 +2327,7 @@ func newArchiveHarness(t *testing.T, now time.Time) *archiveHarness {
db: db,
rawDB: rawDB,
logger: logger,
deps: archiveTestDeps(ctx, t, db),
deps: archiveTestDeps(t, db),
}
}
@@ -2370,16 +2336,13 @@ func newArchiveHarness(t *testing.T, now time.Time) *archiveHarness {
// digest contents.
func createArchiveChat(ctx context.Context, t *testing.T, db database.Store, rawDB *sql.DB, deps chatAutoArchiveDeps, title string, createdAt time.Time) database.Chat {
t.Helper()
chat, err := db.InsertChat(ctx, database.InsertChatParams{
chat := dbgen.Chat(t, db, database.Chat{
OrganizationID: deps.org.ID,
OwnerID: deps.user.ID,
LastModelConfigID: deps.modelConfig.ID,
Title: title,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
})
require.NoError(t, err)
_, err = rawDB.ExecContext(ctx, "UPDATE chats SET created_at = $1, updated_at = $1 WHERE id = $2", createdAt, chat.ID)
_, err := rawDB.ExecContext(ctx, "UPDATE chats SET created_at = $1, updated_at = $1 WHERE id = $2", createdAt, chat.ID)
require.NoError(t, err)
return chat
}
@@ -2389,29 +2352,13 @@ func createArchiveChat(ctx context.Context, t *testing.T, db database.Store, raw
// auto-archive query's LATERAL subquery.
func insertTextMessage(ctx context.Context, t *testing.T, db database.Store, rawDB *sql.DB, chatID, userID, modelConfigID uuid.UUID, createdAt time.Time) {
t.Helper()
msgs, err := db.InsertChatMessages(ctx, database.InsertChatMessagesParams{
ChatID: chatID,
CreatedBy: []uuid.UUID{userID},
ModelConfigID: []uuid.UUID{modelConfigID},
Role: []database.ChatMessageRole{database.ChatMessageRoleUser},
Content: []string{`[{"type":"text","text":"hello"}]`},
ContentVersion: []int16{0},
Visibility: []database.ChatMessageVisibility{database.ChatMessageVisibilityBoth},
InputTokens: []int64{0},
OutputTokens: []int64{0},
TotalTokens: []int64{0},
ReasoningTokens: []int64{0},
CacheCreationTokens: []int64{0},
CacheReadTokens: []int64{0},
ContextLimit: []int64{0},
Compressed: []bool{false},
TotalCostMicros: []int64{0},
RuntimeMs: []int64{0},
ProviderResponseID: []string{""},
msg := dbgen.ChatMessage(t, db, database.ChatMessage{
ChatID: chatID,
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
Role: database.ChatMessageRoleUser,
})
require.NoError(t, err)
require.Len(t, msgs, 1)
_, err = rawDB.ExecContext(ctx, "UPDATE chat_messages SET created_at = $1 WHERE id = $2", createdAt, msgs[0].ID)
_, err := rawDB.ExecContext(ctx, "UPDATE chat_messages SET created_at = $1 WHERE id = $2", createdAt, msg.ID)
require.NoError(t, err)
}
+7 -27
View File
@@ -1260,54 +1260,37 @@ func TestGetAuthorizedChats(t *testing.T) {
dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: secondMember.ID, OrganizationID: org.ID, Roles: []string{rbac.RoleAgentsAccess()}})
// Create FK dependencies: a chat provider and model config.
ctx := testutil.Context(t, testutil.WaitMedium)
_, err = db.InsertChatProvider(ctx, database.InsertChatProviderParams{
Provider: "openai",
DisplayName: "OpenAI",
APIKey: "test-key",
Enabled: true,
CentralApiKeyEnabled: true,
_ = dbgen.ChatProvider(t, db, database.ChatProvider{
Provider: "openai",
DisplayName: "OpenAI",
})
require.NoError(t, err)
modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
Provider: "openai",
Model: "test-model",
DisplayName: "Test Model",
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
Enabled: true,
IsDefault: true,
ContextLimit: 128000,
CompressionThreshold: 80,
Options: json.RawMessage(`{}`),
})
require.NoError(t, err)
// Create 3 chats owned by owner.
for i := range 3 {
_, err := db.InsertChat(ctx, database.InsertChatParams{
dbgen.Chat(t, db, database.Chat{
OrganizationID: org.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: owner.ID,
LastModelConfigID: modelCfg.ID,
Title: fmt.Sprintf("owner chat %d", i+1),
})
require.NoError(t, err)
}
// Create 2 chats owned by member.
for i := range 2 {
_, err := db.InsertChat(ctx, database.InsertChatParams{
dbgen.Chat(t, db, database.Chat{
OrganizationID: org.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: member.ID,
LastModelConfigID: modelCfg.ID,
Title: fmt.Sprintf("member chat %d", i+1),
})
require.NoError(t, err)
}
t.Run("sqlQuerier", func(t *testing.T) {
@@ -1437,15 +1420,12 @@ func TestGetAuthorizedChats(t *testing.T) {
paginationUser := dbgen.User(t, db, database.User{})
dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: paginationUser.ID, OrganizationID: org.ID, Roles: []string{rbac.RoleAgentsAccess()}})
for i := range 7 {
_, err := db.InsertChat(ctx, database.InsertChatParams{
dbgen.Chat(t, db, database.Chat{
OrganizationID: org.ID,
Status: database.ChatStatusWaiting,
ClientType: database.ChatClientTypeUi,
OwnerID: paginationUser.ID,
LastModelConfigID: modelCfg.ID,
Title: fmt.Sprintf("pagination chat %d", i+1),
})
require.NoError(t, err)
}
pagUserSubject, _, err := httpmw.UserRBACSubject(ctx, db, paginationUser.ID, rbac.ExpandableScope(rbac.ScopeAll))