refactor: simplify dataset initialization
This commit is contained in:
@@ -1,20 +1,15 @@
|
|||||||
package initdb
|
package initdb
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
type defaultDatasetSource struct {
|
type defaultDatasetSource struct {
|
||||||
Key string
|
Key string
|
||||||
Name string
|
Name string
|
||||||
URL string
|
URL string
|
||||||
IconURL string
|
IconURL string
|
||||||
LegacyURLs []string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var defaultDatasetSources = []defaultDatasetSource{
|
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-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-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: "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"},
|
||||||
Key: "rsshub-cls-hot", Name: "财联社-热门",
|
{Key: "rsshub-jin10", Name: "金十数据-快讯", URL: "https://rsshub.ktachibana.party/jin10"},
|
||||||
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"},
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// InitDataset inserts or repairs root's built-in dataset sources without touching user-created sources.
|
// InitDataset inserts the built-in dataset sources when the owner has no sources.
|
||||||
func InitDataset(database *gorm.DB) error {
|
func InitDataset(database *gorm.DB, ownerID uint, ownerIdentity string) error {
|
||||||
var root models.SaUser
|
var count int64
|
||||||
if err := database.Where("email = ? AND role = ?", rootUsername, rootRole).First(&root).Error; err != nil {
|
if err := database.Model(&models.SaDatasetSource{}).
|
||||||
return fmt.Errorf("find root user for default dataset sources: %w", err)
|
Where("owner_id = ?", ownerID).
|
||||||
|
Count(&count).Error; err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
if count > 0 {
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
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
|
for _, definition := range defaultDatasetSources {
|
||||||
if err := tx.Where("owner_id = ? OR owner_identity = ?", 0, "").Find(&sources).Error; err != nil {
|
seedKey := definition.Key
|
||||||
return err
|
if err := database.Create(&models.SaDatasetSource{
|
||||||
}
|
OwnerID: ownerID,
|
||||||
for _, source := range sources {
|
OwnerIdentity: ownerIdentity,
|
||||||
var owner models.SaUser
|
SeedKey: &seedKey,
|
||||||
if err := tx.Select("identity").First(&owner, source.OwnerID).Error; err != nil {
|
Name: definition.Name,
|
||||||
return err
|
Kind: "rss",
|
||||||
}
|
URL: definition.URL,
|
||||||
if err := tx.Model(&source).Updates(map[string]any{
|
IconURL: definition.IconURL,
|
||||||
"owner_identity": owner.Identity,
|
Description: "默认资讯订阅",
|
||||||
|
Enabled: true,
|
||||||
}).Error; err != nil {
|
}).Error; err != nil {
|
||||||
return err
|
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
|
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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ package initdb
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -11,19 +10,20 @@ import (
|
|||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestInitDatasetCreatesDefaultSourcesWhenTableIsEmpty(t *testing.T) {
|
func TestInitDatasetCreatesDefaultSourcesWhenOwnerHasNoSources(t *testing.T) {
|
||||||
database := newDatasetInitTestDatabase(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, root.ID, root.Identity))
|
||||||
require.NoError(t, InitDataset(database))
|
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
|
var sources []models.SaDatasetSource
|
||||||
require.NoError(t, database.Order("id asc").Find(&sources).Error)
|
require.NoError(t, database.Order("id asc").Find(&sources).Error)
|
||||||
require.Len(t, sources, len(defaultDatasetSources))
|
require.Len(t, sources, len(defaultDatasetSources))
|
||||||
for index, expected := range 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.Name, sources[index].Name)
|
||||||
require.Equal(t, expected.URL, sources[index].URL)
|
require.Equal(t, expected.URL, sources[index].URL)
|
||||||
require.Equal(t, expected.IconURL, sources[index].IconURL)
|
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, expected.Key, *sources[index].SeedKey)
|
||||||
require.Equal(t, "rss", sources[index].Kind)
|
require.Equal(t, "rss", sources[index].Kind)
|
||||||
require.True(t, sources[index].Enabled)
|
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)
|
database := newDatasetInitTestDatabase(t)
|
||||||
require.NoError(t, InitUser(database))
|
root, err := InitUser(database)
|
||||||
var root models.SaUser
|
require.NoError(t, err)
|
||||||
require.NoError(t, database.Where("email = ?", rootUsername).First(&root).Error)
|
|
||||||
existing := models.SaDatasetSource{
|
existing := models.SaDatasetSource{
|
||||||
OwnerID: root.ID,
|
OwnerID: root.ID,
|
||||||
Name: "Existing source",
|
Name: "Existing source",
|
||||||
@@ -49,151 +46,51 @@ func TestInitDatasetAddsMissingDefaultsWithoutChangingCustomSources(t *testing.T
|
|||||||
}
|
}
|
||||||
require.NoError(t, database.Create(&existing).Error)
|
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
|
var sources []models.SaDatasetSource
|
||||||
require.NoError(t, database.Find(&sources).Error)
|
require.NoError(t, database.Find(&sources).Error)
|
||||||
require.Len(t, sources, len(defaultDatasetSources)+1)
|
require.Len(t, sources, 1)
|
||||||
var preserved models.SaDatasetSource
|
require.Equal(t, existing.Identity, sources[0].Identity)
|
||||||
require.NoError(t, database.Where("identity = ?", existing.Identity).First(&preserved).Error)
|
|
||||||
require.Equal(t, existing.Name, preserved.Name)
|
|
||||||
require.Nil(t, preserved.SeedKey)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInitDatasetBackfillsLegacySourceOwners(t *testing.T) {
|
func TestInitDatasetCountsSourcesByOwner(t *testing.T) {
|
||||||
database := newDatasetInitTestDatabase(t)
|
database := newDatasetInitTestDatabase(t)
|
||||||
require.NoError(t, InitUser(database))
|
root, err := InitUser(database)
|
||||||
|
require.NoError(t, err)
|
||||||
other := models.SaUser{
|
other := models.SaUser{
|
||||||
Email: "other@example.com", DisplayName: "Other",
|
Email: "other@example.com", DisplayName: "Other",
|
||||||
PasswordHash: "hash", Role: "user",
|
PasswordHash: "hash", Role: "user",
|
||||||
}
|
}
|
||||||
require.NoError(t, database.Create(&other).Error)
|
require.NoError(t, database.Create(&other).Error)
|
||||||
require.NoError(t, database.Exec(
|
require.NoError(t, database.Create(&models.SaDatasetSource{
|
||||||
"ALTER TABLE sa_dataset_sources ADD COLUMN created_by integer NOT NULL DEFAULT 0",
|
OwnerID: other.ID, Name: "Other source", Kind: "manual", Enabled: true,
|
||||||
).Error)
|
}).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, InitDataset(database))
|
require.NoError(t, InitDataset(database, root.ID, root.Identity))
|
||||||
|
|
||||||
require.NoError(t, database.First(&source, source.ID).Error)
|
var rootCount int64
|
||||||
require.Equal(t, other.ID, source.OwnerID)
|
require.NoError(t, database.Model(&models.SaDatasetSource{}).
|
||||||
require.Equal(t, other.Identity, source.OwnerIdentity)
|
Where("owner_id = ?", root.ID).
|
||||||
hasCreatedBy, err := hasDatasetSourceColumn(database, "created_by")
|
Count(&rootCount).Error)
|
||||||
require.NoError(t, err)
|
require.Equal(t, int64(len(defaultDatasetSources)), rootCount)
|
||||||
require.False(t, hasCreatedBy)
|
var otherCount int64
|
||||||
hasCreatedByIdentity, err := hasDatasetSourceColumn(database, "created_by_identity")
|
require.NoError(t, database.Model(&models.SaDatasetSource{}).
|
||||||
require.NoError(t, err)
|
Where("owner_id = ?", other.ID).
|
||||||
require.False(t, hasCreatedByIdentity)
|
Count(&otherCount).Error)
|
||||||
}
|
require.Equal(t, int64(1), otherCount)
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newDatasetInitTestDatabase(t *testing.T) *gorm.DB {
|
func newDatasetInitTestDatabase(t *testing.T) *gorm.DB {
|
||||||
t.Helper()
|
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, err)
|
||||||
require.NoError(t, database.AutoMigrate(
|
require.NoError(t, database.AutoMigrate(
|
||||||
&models.SaUser{},
|
&models.SaUser{},
|
||||||
&models.SaDatasetSource{},
|
&models.SaDatasetSource{},
|
||||||
&models.SaDatasetItem{},
|
|
||||||
))
|
))
|
||||||
return database
|
return database
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,10 +5,11 @@ import "gorm.io/gorm"
|
|||||||
// New initializes the default database records.
|
// New initializes the default database records.
|
||||||
func New(database *gorm.DB) error {
|
func New(database *gorm.DB) error {
|
||||||
return database.Transaction(func(tx *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
|
return err
|
||||||
}
|
}
|
||||||
if err := InitDataset(tx); err != nil {
|
if err := InitDataset(tx, user.ID, user.Identity); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return InitExpert(tx)
|
return InitExpert(tx)
|
||||||
|
|||||||
@@ -12,25 +12,34 @@ const (
|
|||||||
rootRole = "root"
|
rootRole = "root"
|
||||||
)
|
)
|
||||||
|
|
||||||
// InitUser creates the default root account when the user table is empty.
|
// InitUser returns the existing root user or creates and returns it when the user table is empty.
|
||||||
func InitUser(database *gorm.DB) error {
|
func InitUser(database *gorm.DB) (*models.SaUser, error) {
|
||||||
var count int64
|
var count int64
|
||||||
if err := database.Model(&models.SaUser{}).Count(&count).Error; err != nil {
|
if err := database.Model(&models.SaUser{}).Count(&count).Error; err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
if count > 0 {
|
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)
|
passwordHash, err := bcrypt.GenerateFromPassword([]byte(rootPassword), bcrypt.DefaultCost)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return database.Create(&models.SaUser{
|
user := &models.SaUser{
|
||||||
Email: rootUsername,
|
Email: rootUsername,
|
||||||
DisplayName: rootUsername,
|
DisplayName: rootUsername,
|
||||||
PasswordHash: string(passwordHash),
|
PasswordHash: string(passwordHash),
|
||||||
Role: rootRole,
|
Role: rootRole,
|
||||||
}).Error
|
}
|
||||||
|
if err := database.Create(user).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -13,33 +13,42 @@ import (
|
|||||||
func TestInitUserCreatesRootUserWhenTableIsEmpty(t *testing.T) {
|
func TestInitUserCreatesRootUserWhenTableIsEmpty(t *testing.T) {
|
||||||
database := newUserTestDatabase(t)
|
database := newUserTestDatabase(t)
|
||||||
|
|
||||||
require.NoError(t, InitUser(database))
|
user, err := InitUser(database)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
var users []models.SaUser
|
var users []models.SaUser
|
||||||
require.NoError(t, database.Find(&users).Error)
|
require.NoError(t, database.Find(&users).Error)
|
||||||
require.Len(t, users, 1)
|
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].Email)
|
||||||
require.Equal(t, rootUsername, users[0].DisplayName)
|
require.Equal(t, rootUsername, users[0].DisplayName)
|
||||||
require.Equal(t, rootRole, users[0].Role)
|
require.Equal(t, rootRole, users[0].Role)
|
||||||
require.NoError(t, bcrypt.CompareHashAndPassword([]byte(users[0].PasswordHash), []byte(rootPassword)))
|
require.NoError(t, bcrypt.CompareHashAndPassword([]byte(users[0].PasswordHash), []byte(rootPassword)))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInitUserDoesNothingWhenTableIsNotEmpty(t *testing.T) {
|
func TestInitUserReturnsExistingRootWhenTableIsNotEmpty(t *testing.T) {
|
||||||
database := newUserTestDatabase(t)
|
database := newUserTestDatabase(t)
|
||||||
existing := models.SaUser{
|
root := models.SaUser{
|
||||||
Email: "existing@example.com",
|
Email: rootUsername,
|
||||||
DisplayName: "Existing User",
|
DisplayName: rootUsername,
|
||||||
PasswordHash: "existing-hash",
|
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
|
var users []models.SaUser
|
||||||
require.NoError(t, database.Find(&users).Error)
|
require.NoError(t, database.Find(&users).Error)
|
||||||
require.Len(t, users, 1)
|
require.Len(t, users, 2)
|
||||||
require.Equal(t, existing.Email, users[0].Email)
|
require.Equal(t, root.ID, user.ID)
|
||||||
|
require.Equal(t, root.Identity, user.Identity)
|
||||||
}
|
}
|
||||||
|
|
||||||
func newUserTestDatabase(t *testing.T) *gorm.DB {
|
func newUserTestDatabase(t *testing.T) *gorm.DB {
|
||||||
|
|||||||
Reference in New Issue
Block a user