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 {