mirror of
https://github.com/coder/coder.git
synced 2026-09-22 05:05:20 +08:00
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:
@@ -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(¶ms)
|
||||
}
|
||||
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(¶ms)
|
||||
}
|
||||
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()),
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user