diff --git a/backend/internal/initdb/dataset.go b/backend/internal/initdb/dataset.go index 6305528..cd9ba42 100644 --- a/backend/internal/initdb/dataset.go +++ b/backend/internal/initdb/dataset.go @@ -1,20 +1,15 @@ package initdb import ( - "errors" - "fmt" - "strings" - "gorm.io/gorm" "senlinai-agent/backend/internal/models" ) type defaultDatasetSource struct { - Key string - Name string - URL string - IconURL string - LegacyURLs []string + Key string + Name string + URL string + IconURL string } var defaultDatasetSources = []defaultDatasetSource{ @@ -24,179 +19,37 @@ var defaultDatasetSources = []defaultDatasetSource{ {Key: "buzzing-wsj", Name: "华尔街日报热门", URL: "https://wsj.buzzing.cc/feed.json", IconURL: "https://wsj.buzzing.cc/apple-touch-icon.png"}, {Key: "buzzing-product-hunt", Name: "Product Hunt", URL: "https://ph.buzzing.cc/feed.json", IconURL: "https://ph.buzzing.cc/apple-touch-icon.png"}, {Key: "buzzing-devto", Name: "Dev.to", URL: "https://dev.buzzing.cc/feed.json", IconURL: "https://dev.buzzing.cc/apple-touch-icon.png"}, - { - Key: "rsshub-cls-hot", Name: "财联社-热门", - URL: "https://rsshub.ktachibana.party/cls/hot", - LegacyURLs: []string{"https://rsshub.app/cls/hot"}, - }, - { - Key: "rsshub-jin10", Name: "金十数据-快讯", - URL: "https://rsshub.ktachibana.party/jin10", - LegacyURLs: []string{"https://rsshub.app/jin10"}, - }, + {Key: "rsshub-cls-hot", Name: "财联社-热门", URL: "https://rsshub.ktachibana.party/cls/hot"}, + {Key: "rsshub-jin10", Name: "金十数据-快讯", URL: "https://rsshub.ktachibana.party/jin10"}, } -// InitDataset inserts or repairs root's built-in dataset sources without touching user-created sources. -func InitDataset(database *gorm.DB) error { - var root models.SaUser - if err := database.Where("email = ? AND role = ?", rootUsername, rootRole).First(&root).Error; err != nil { - return fmt.Errorf("find root user for default dataset sources: %w", err) +// InitDataset inserts the built-in dataset sources when the owner has no sources. +func InitDataset(database *gorm.DB, ownerID uint, ownerIdentity string) error { + var count int64 + if err := database.Model(&models.SaDatasetSource{}). + Where("owner_id = ?", ownerID). + Count(&count).Error; err != nil { + return err } - - return database.Transaction(func(tx *gorm.DB) error { - if err := prepareDatasetSourceOwners(tx); err != nil { - return err - } - if err := prepareDatasetItems(tx); err != nil { - return err - } - for _, definition := range defaultDatasetSources { - if err := upsertDefaultDatasetSource(tx, root, definition); err != nil { - return err - } - } + if count > 0 { return nil - }) -} - -func prepareDatasetSourceOwners(tx *gorm.DB) error { - hasCreatedBy, err := hasDatasetSourceColumn(tx, "created_by") - if err != nil { - return err - } - if hasCreatedBy { - if err := tx.Exec( - "UPDATE sa_dataset_sources SET owner_id = created_by WHERE owner_id = 0", - ).Error; err != nil { - return fmt.Errorf("backfill dataset source owners: %w", err) - } } - var sources []models.SaDatasetSource - if err := tx.Where("owner_id = ? OR owner_identity = ?", 0, "").Find(&sources).Error; err != nil { - return err - } - for _, source := range sources { - var owner models.SaUser - if err := tx.Select("identity").First(&owner, source.OwnerID).Error; err != nil { - return err - } - if err := tx.Model(&source).Updates(map[string]any{ - "owner_identity": owner.Identity, + for _, definition := range defaultDatasetSources { + seedKey := definition.Key + if err := database.Create(&models.SaDatasetSource{ + OwnerID: ownerID, + OwnerIdentity: ownerIdentity, + SeedKey: &seedKey, + Name: definition.Name, + Kind: "rss", + URL: definition.URL, + IconURL: definition.IconURL, + Description: "默认资讯订阅", + Enabled: true, }).Error; err != nil { return err } } - - if tx.Migrator().HasIndex(&models.SaDatasetSource{}, "idx_sa_dataset_source_seed") { - if err := tx.Migrator().DropIndex(&models.SaDatasetSource{}, "idx_sa_dataset_source_seed"); err != nil { - return fmt.Errorf("drop legacy dataset source seed index: %w", err) - } - } - for _, index := range []string{ - "idx_sa_dataset_sources_created_by", - "idx_sa_dataset_sources_created_by_identity", - } { - if tx.Migrator().HasIndex(&models.SaDatasetSource{}, index) { - if err := tx.Migrator().DropIndex(&models.SaDatasetSource{}, index); err != nil { - return fmt.Errorf("drop legacy dataset source creator index %s: %w", index, err) - } - } - } - hasCreatedByIdentity, err := hasDatasetSourceColumn(tx, "created_by_identity") - if err != nil { - return err - } - if hasCreatedByIdentity { - if err := tx.Exec( - "ALTER TABLE sa_dataset_sources DROP COLUMN created_by_identity", - ).Error; err != nil { - return fmt.Errorf("drop legacy dataset source creator identity: %w", err) - } - } - hasCreatedBy, err = hasDatasetSourceColumn(tx, "created_by") - if err != nil { - return err - } - if hasCreatedBy { - if err := tx.Exec( - "ALTER TABLE sa_dataset_sources DROP COLUMN created_by", - ).Error; err != nil { - return fmt.Errorf("drop legacy dataset source creator: %w", err) - } - } return nil } - -func hasDatasetSourceColumn(tx *gorm.DB, name string) (bool, error) { - columns, err := tx.Migrator().ColumnTypes(&models.SaDatasetSource{}) - if err != nil { - return false, fmt.Errorf("inspect dataset source columns: %w", err) - } - for _, column := range columns { - if strings.EqualFold(column.Name(), name) { - return true, nil - } - } - return false, nil -} - -func prepareDatasetItems(tx *gorm.DB) error { - if tx.Migrator().HasIndex(&models.SaDatasetItem{}, "idx_sa_dataset_item_source_external") { - return nil - } - if err := tx.Exec( - "UPDATE sa_dataset_items SET external_id = NULL WHERE external_id = ''", - ).Error; err != nil { - return err - } - if err := tx.Exec(` - DELETE FROM sa_dataset_items - WHERE external_id IS NOT NULL - AND EXISTS ( - SELECT 1 - FROM sa_dataset_items AS earlier - WHERE earlier.source_id = sa_dataset_items.source_id - AND earlier.external_id = sa_dataset_items.external_id - AND earlier.id < sa_dataset_items.id - ) - `).Error; err != nil { - return err - } - return tx.Exec(` - CREATE UNIQUE INDEX IF NOT EXISTS idx_sa_dataset_item_source_external - ON sa_dataset_items (source_id, external_id) - `).Error -} - -func upsertDefaultDatasetSource(tx *gorm.DB, root models.SaUser, definition defaultDatasetSource) error { - var source models.SaDatasetSource - err := tx.Where("owner_id = ? AND seed_key = ?", root.ID, definition.Key).First(&source).Error - if err == nil { - return tx.Model(&source).Updates(map[string]any{ - "url": definition.URL, "icon_url": definition.IconURL, "kind": "rss", - }).Error - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return err - } - - knownURLs := append([]string{definition.URL}, definition.LegacyURLs...) - err = tx.Where("owner_id = ? AND url IN ?", root.ID, knownURLs).First(&source).Error - if err == nil { - return tx.Model(&source).Updates(map[string]any{ - "seed_key": definition.Key, "url": definition.URL, - "icon_url": definition.IconURL, "kind": "rss", - }).Error - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return err - } - - seedKey := definition.Key - return tx.Create(&models.SaDatasetSource{ - OwnerID: root.ID, SeedKey: &seedKey, Name: definition.Name, - Kind: "rss", URL: definition.URL, IconURL: definition.IconURL, - Description: "默认资讯订阅", Enabled: true, - }).Error -} diff --git a/backend/internal/initdb/dataset_test.go b/backend/internal/initdb/dataset_test.go index d26808e..b1ff313 100644 --- a/backend/internal/initdb/dataset_test.go +++ b/backend/internal/initdb/dataset_test.go @@ -3,7 +3,6 @@ package initdb import ( "fmt" "testing" - "time" "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" @@ -11,19 +10,20 @@ import ( "senlinai-agent/backend/internal/models" ) -func TestInitDatasetCreatesDefaultSourcesWhenTableIsEmpty(t *testing.T) { +func TestInitDatasetCreatesDefaultSourcesWhenOwnerHasNoSources(t *testing.T) { database := newDatasetInitTestDatabase(t) - require.NoError(t, InitUser(database)) + root, err := InitUser(database) + require.NoError(t, err) - require.NoError(t, InitDataset(database)) - require.NoError(t, InitDataset(database)) + require.NoError(t, InitDataset(database, root.ID, root.Identity)) + require.NoError(t, InitDataset(database, root.ID, root.Identity)) - var root models.SaUser - require.NoError(t, database.Where("email = ?", rootUsername).First(&root).Error) var sources []models.SaDatasetSource require.NoError(t, database.Order("id asc").Find(&sources).Error) require.Len(t, sources, len(defaultDatasetSources)) for index, expected := range defaultDatasetSources { + require.Equal(t, root.ID, sources[index].OwnerID) + require.Equal(t, root.Identity, sources[index].OwnerIdentity) require.Equal(t, expected.Name, sources[index].Name) require.Equal(t, expected.URL, sources[index].URL) require.Equal(t, expected.IconURL, sources[index].IconURL) @@ -31,16 +31,13 @@ func TestInitDatasetCreatesDefaultSourcesWhenTableIsEmpty(t *testing.T) { require.Equal(t, expected.Key, *sources[index].SeedKey) require.Equal(t, "rss", sources[index].Kind) require.True(t, sources[index].Enabled) - require.Equal(t, root.ID, sources[index].OwnerID) - require.Equal(t, root.Identity, sources[index].OwnerIdentity) } } -func TestInitDatasetAddsMissingDefaultsWithoutChangingCustomSources(t *testing.T) { +func TestInitDatasetDoesNothingWhenOwnerAlreadyHasSource(t *testing.T) { database := newDatasetInitTestDatabase(t) - require.NoError(t, InitUser(database)) - var root models.SaUser - require.NoError(t, database.Where("email = ?", rootUsername).First(&root).Error) + root, err := InitUser(database) + require.NoError(t, err) existing := models.SaDatasetSource{ OwnerID: root.ID, Name: "Existing source", @@ -49,151 +46,51 @@ func TestInitDatasetAddsMissingDefaultsWithoutChangingCustomSources(t *testing.T } require.NoError(t, database.Create(&existing).Error) - require.NoError(t, InitDataset(database)) + require.NoError(t, InitDataset(database, root.ID, root.Identity)) var sources []models.SaDatasetSource require.NoError(t, database.Find(&sources).Error) - require.Len(t, sources, len(defaultDatasetSources)+1) - var preserved models.SaDatasetSource - require.NoError(t, database.Where("identity = ?", existing.Identity).First(&preserved).Error) - require.Equal(t, existing.Name, preserved.Name) - require.Nil(t, preserved.SeedKey) + require.Len(t, sources, 1) + require.Equal(t, existing.Identity, sources[0].Identity) } -func TestInitDatasetBackfillsLegacySourceOwners(t *testing.T) { +func TestInitDatasetCountsSourcesByOwner(t *testing.T) { database := newDatasetInitTestDatabase(t) - require.NoError(t, InitUser(database)) + root, err := InitUser(database) + require.NoError(t, err) other := models.SaUser{ Email: "other@example.com", DisplayName: "Other", PasswordHash: "hash", Role: "user", } require.NoError(t, database.Create(&other).Error) - require.NoError(t, database.Exec( - "ALTER TABLE sa_dataset_sources ADD COLUMN created_by integer NOT NULL DEFAULT 0", - ).Error) - require.NoError(t, database.Exec( - "ALTER TABLE sa_dataset_sources ADD COLUMN created_by_identity text", - ).Error) - source := models.SaDatasetSource{ - OwnerID: other.ID, Name: "Existing source", Kind: "manual", Enabled: true, - } - require.NoError(t, database.Create(&source).Error) - require.NoError(t, database.Exec(` - UPDATE sa_dataset_sources - SET owner_id = 0, owner_identity = '', created_by = ?, created_by_identity = ? - WHERE id = ? - `, other.ID, other.Identity, source.ID).Error) + require.NoError(t, database.Create(&models.SaDatasetSource{ + OwnerID: other.ID, Name: "Other source", Kind: "manual", Enabled: true, + }).Error) - require.NoError(t, InitDataset(database)) + require.NoError(t, InitDataset(database, root.ID, root.Identity)) - require.NoError(t, database.First(&source, source.ID).Error) - require.Equal(t, other.ID, source.OwnerID) - require.Equal(t, other.Identity, source.OwnerIdentity) - hasCreatedBy, err := hasDatasetSourceColumn(database, "created_by") - require.NoError(t, err) - require.False(t, hasCreatedBy) - hasCreatedByIdentity, err := hasDatasetSourceColumn(database, "created_by_identity") - require.NoError(t, err) - require.False(t, hasCreatedByIdentity) -} - -func TestInitDatasetRejectsMissingRootInsteadOfAssigningDefaultsToAnotherUser(t *testing.T) { - database := newDatasetInitTestDatabase(t) - user := models.SaUser{ - Email: "existing@example.com", - DisplayName: "Existing", - PasswordHash: "hash", - Role: "user", - } - require.NoError(t, database.Create(&user).Error) - - require.Error(t, InitDataset(database)) - - var sources []models.SaDatasetSource - require.NoError(t, database.Find(&sources).Error) - require.Empty(t, sources) -} - -func TestInitDatasetRepairsLegacyRSSHubURLsForRootOnly(t *testing.T) { - database := newDatasetInitTestDatabase(t) - require.NoError(t, InitUser(database)) - var root models.SaUser - require.NoError(t, database.Where("email = ?", rootUsername).First(&root).Error) - other := models.SaUser{ - Email: "other@example.com", DisplayName: "Other", - PasswordHash: "hash", Role: "user", - } - require.NoError(t, database.Create(&other).Error) - - rootLegacy := models.SaDatasetSource{ - OwnerID: root.ID, Name: "财联社-热门", Kind: "rss", - URL: "https://rsshub.app/cls/hot", Enabled: true, - } - otherLegacy := models.SaDatasetSource{ - OwnerID: other.ID, Name: "财联社-热门", Kind: "rss", - URL: "https://rsshub.app/cls/hot", Enabled: true, - } - require.NoError(t, database.Create(&rootLegacy).Error) - require.NoError(t, database.Create(&otherLegacy).Error) - - require.NoError(t, InitDataset(database)) - - require.NoError(t, database.First(&rootLegacy, rootLegacy.ID).Error) - require.Equal(t, "https://rsshub.ktachibana.party/cls/hot", rootLegacy.URL) - require.NotNil(t, rootLegacy.SeedKey) - require.Equal(t, "rsshub-cls-hot", *rootLegacy.SeedKey) - require.NoError(t, database.First(&otherLegacy, otherLegacy.ID).Error) - require.Equal(t, "https://rsshub.app/cls/hot", otherLegacy.URL) - require.Nil(t, otherLegacy.SeedKey) -} - -func TestInitDatasetPreparesCollectedItemUniqueness(t *testing.T) { - database := newDatasetInitTestDatabase(t) - require.NoError(t, InitUser(database)) - var root models.SaUser - require.NoError(t, database.Where("email = ?", rootUsername).First(&root).Error) - source := models.SaDatasetSource{ - OwnerID: root.ID, Name: "Existing feed", Kind: "rss", - URL: "https://example.com/feed", Enabled: true, - } - require.NoError(t, database.Create(&source).Error) - now := time.Now() - for index, externalID := range []string{"duplicate", "duplicate", "", ""} { - require.NoError(t, database.Exec(` - INSERT INTO sa_dataset_items - (identity, source_id, created_by, external_id, title, status, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?) - `, fmt.Sprintf("item-%d", index), source.ID, root.ID, externalID, - fmt.Sprintf("Item %d", index), "unread", now, now).Error) - } - - require.NoError(t, InitDataset(database)) - - var collectedCount int64 - require.NoError(t, database.Model(&models.SaDatasetItem{}). - Where("source_id = ? AND external_id = ?", source.ID, "duplicate"). - Count(&collectedCount).Error) - require.Equal(t, int64(1), collectedCount) - var manualCount int64 - require.NoError(t, database.Model(&models.SaDatasetItem{}). - Where("source_id = ? AND external_id IS NULL", source.ID). - Count(&manualCount).Error) - require.Equal(t, int64(2), manualCount) - require.Error(t, database.Exec(` - INSERT INTO sa_dataset_items - (identity, source_id, created_by, external_id, title, status, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?) - `, "item-new", source.ID, root.ID, "duplicate", "Duplicate", "unread", now, now).Error) + var rootCount int64 + require.NoError(t, database.Model(&models.SaDatasetSource{}). + Where("owner_id = ?", root.ID). + Count(&rootCount).Error) + require.Equal(t, int64(len(defaultDatasetSources)), rootCount) + var otherCount int64 + require.NoError(t, database.Model(&models.SaDatasetSource{}). + Where("owner_id = ?", other.ID). + Count(&otherCount).Error) + require.Equal(t, int64(1), otherCount) } func newDatasetInitTestDatabase(t *testing.T) *gorm.DB { t.Helper() - database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{}) + database, err := gorm.Open( + sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), + &gorm.Config{}, + ) require.NoError(t, err) require.NoError(t, database.AutoMigrate( &models.SaUser{}, &models.SaDatasetSource{}, - &models.SaDatasetItem{}, )) return database } diff --git a/backend/internal/initdb/new.go b/backend/internal/initdb/new.go index ed9808e..9a71030 100644 --- a/backend/internal/initdb/new.go +++ b/backend/internal/initdb/new.go @@ -5,10 +5,11 @@ import "gorm.io/gorm" // New initializes the default database records. func New(database *gorm.DB) error { return database.Transaction(func(tx *gorm.DB) error { - if err := InitUser(tx); err != nil { + user, err := InitUser(tx) + if err != nil { return err } - if err := InitDataset(tx); err != nil { + if err := InitDataset(tx, user.ID, user.Identity); err != nil { return err } return InitExpert(tx) diff --git a/backend/internal/initdb/user.go b/backend/internal/initdb/user.go index 02648d4..5dc81fe 100644 --- a/backend/internal/initdb/user.go +++ b/backend/internal/initdb/user.go @@ -12,25 +12,34 @@ const ( rootRole = "root" ) -// InitUser creates the default root account when the user table is empty. -func InitUser(database *gorm.DB) error { +// InitUser returns the existing root user or creates and returns it when the user table is empty. +func InitUser(database *gorm.DB) (*models.SaUser, error) { var count int64 if err := database.Model(&models.SaUser{}).Count(&count).Error; err != nil { - return err + return nil, err } if count > 0 { - return nil + var user models.SaUser + if err := database.Where("email = ? AND role = ?", rootUsername, rootRole). + First(&user).Error; err != nil { + return nil, err + } + return &user, nil } passwordHash, err := bcrypt.GenerateFromPassword([]byte(rootPassword), bcrypt.DefaultCost) if err != nil { - return err + return nil, err } - return database.Create(&models.SaUser{ + user := &models.SaUser{ Email: rootUsername, DisplayName: rootUsername, PasswordHash: string(passwordHash), Role: rootRole, - }).Error + } + if err := database.Create(user).Error; err != nil { + return nil, err + } + return user, nil } diff --git a/backend/internal/initdb/user_test.go b/backend/internal/initdb/user_test.go index c5a7541..c0e95e2 100644 --- a/backend/internal/initdb/user_test.go +++ b/backend/internal/initdb/user_test.go @@ -13,33 +13,42 @@ import ( func TestInitUserCreatesRootUserWhenTableIsEmpty(t *testing.T) { database := newUserTestDatabase(t) - require.NoError(t, InitUser(database)) + user, err := InitUser(database) + require.NoError(t, err) var users []models.SaUser require.NoError(t, database.Find(&users).Error) require.Len(t, users, 1) + require.Equal(t, users[0].ID, user.ID) + require.Equal(t, users[0].Identity, user.Identity) require.Equal(t, rootUsername, users[0].Email) require.Equal(t, rootUsername, users[0].DisplayName) require.Equal(t, rootRole, users[0].Role) require.NoError(t, bcrypt.CompareHashAndPassword([]byte(users[0].PasswordHash), []byte(rootPassword))) } -func TestInitUserDoesNothingWhenTableIsNotEmpty(t *testing.T) { +func TestInitUserReturnsExistingRootWhenTableIsNotEmpty(t *testing.T) { database := newUserTestDatabase(t) - existing := models.SaUser{ - Email: "existing@example.com", - DisplayName: "Existing User", + root := models.SaUser{ + Email: rootUsername, + DisplayName: rootUsername, PasswordHash: "existing-hash", - Role: "user", + Role: rootRole, } - require.NoError(t, database.Create(&existing).Error) + require.NoError(t, database.Create(&root).Error) + require.NoError(t, database.Create(&models.SaUser{ + Email: "other@example.com", DisplayName: "Other", + PasswordHash: "hash", Role: "user", + }).Error) - require.NoError(t, InitUser(database)) + user, err := InitUser(database) + require.NoError(t, err) var users []models.SaUser require.NoError(t, database.Find(&users).Error) - require.Len(t, users, 1) - require.Equal(t, existing.Email, users[0].Email) + require.Len(t, users, 2) + require.Equal(t, root.ID, user.ID) + require.Equal(t, root.Identity, user.Identity) } func newUserTestDatabase(t *testing.T) *gorm.DB {