fix: complete exploration data flow

This commit is contained in:
2026-07-23 22:25:12 +08:00
parent 55ff320cec
commit 5a1b04c4ed
12 changed files with 952 additions and 189 deletions

View File

@@ -3,9 +3,11 @@ package dataset
import (
"context"
"errors"
"fmt"
"net/url"
"os"
"testing"
"time"
"github.com/stretchr/testify/require"
"senlinai-agent/backend/internal/models"
@@ -164,8 +166,68 @@ func TestSyncAllSourcesCollectsEnabledRSSSourcesForEveryUser(t *testing.T) {
require.Equal(t, int64(2), itemCount)
}
func TestSyncSourcesCollectsMultipleSourcesForOneOwner(t *testing.T) {
database := newDatasetTestDatabase(t)
user := createDatasetTestUser(t, database, "multi-feed-owner@example.com")
for index := 0; index < 5; index++ {
require.NoError(t, database.Create(&models.SaDatasetSource{
CreatedBy: user.ID, Name: fmt.Sprintf("Feed %d", index), Kind: "rss",
URL: fmt.Sprintf("https://example.com/feed-%d", index), Enabled: true,
}).Error)
}
runs, err := newServiceWithFetcher(database, staticFeedFetcher{}).SyncSources(context.Background(), user.ID)
require.NoError(t, err)
require.Len(t, runs, 5)
for _, run := range runs {
require.Equal(t, "completed", run.Status)
}
}
func TestQueueSourcesRunsInBackgroundAndRejectsOverlap(t *testing.T) {
database := newDatasetTestDatabase(t)
user := createDatasetTestUser(t, database, "queued-feed-owner@example.com")
source := models.SaDatasetSource{
CreatedBy: user.ID, Name: "Queued RSS", Kind: "rss",
URL: "https://example.com/feed", Enabled: true,
}
require.NoError(t, database.Create(&source).Error)
fetcher := &blockingFeedFetcher{
started: make(chan struct{}),
release: make(chan struct{}),
}
service := newServiceWithFetcher(database, fetcher)
queued, err := service.QueueSources(user.ID)
require.NoError(t, err)
require.Len(t, queued, 1)
require.Equal(t, "pending", queued[0].Status)
<-fetcher.started
_, err = service.QueueSources(user.ID)
require.ErrorIs(t, err, ErrSyncInProgress)
require.ErrorIs(t, service.DeleteSource(user.ID, source.Identity), ErrSyncInProgress)
close(fetcher.release)
require.Eventually(t, func() bool {
runs, listErr := service.ListCrons(user.ID)
return listErr == nil && len(runs) == 1 && runs[0].Status == "completed"
}, time.Second, 10*time.Millisecond)
}
type failingFeedFetcher struct{}
func (failingFeedFetcher) Fetch(context.Context, string) (ParsedFeed, error) {
return ParsedFeed{}, errors.New("feed unavailable")
}
type blockingFeedFetcher struct {
started chan struct{}
release chan struct{}
}
func (f *blockingFeedFetcher) Fetch(context.Context, string) (ParsedFeed, error) {
close(f.started)
<-f.release
return staticFeedFetcher{}.Fetch(context.Background(), "")
}

View File

@@ -3,6 +3,7 @@ package dataset
import (
"errors"
"net/http"
"strconv"
"strings"
"time"
@@ -25,6 +26,7 @@ type SourceDTO struct {
IconURL string `json:"iconUrl"`
Description string `json:"description"`
Enabled bool `json:"enabled"`
BuiltIn bool `json:"builtIn"`
LastSyncedAt *time.Time `json:"lastSyncedAt"`
ItemCount int64 `json:"itemCount"`
CreatedAt time.Time `json:"createdAt"`
@@ -45,19 +47,28 @@ type ItemDTO struct {
UpdatedAt time.Time `json:"updatedAt"`
}
type ItemPageDTO struct {
Items []ItemDTO `json:"items"`
Total int64 `json:"total"`
Limit int `json:"limit"`
Offset int `json:"offset"`
}
type CronDTO struct {
ID string `json:"id"`
SourceID string `json:"sourceId"`
Schedule string `json:"schedule"`
Status string `json:"status"`
Enabled bool `json:"enabled"`
NextRunAt *time.Time `json:"nextRunAt"`
LastRunAt *time.Time `json:"lastRunAt"`
LastResult string `json:"lastResult"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
type DepositDTO struct {
NoteID string `json:"noteId"`
ProjectID string `json:"projectId"`
}
type sourceRequest struct {
Name string `json:"name"`
Kind string `json:"kind"`
@@ -79,6 +90,7 @@ func (h *Handler) Register(router gin.IRouter) {
router.GET("/dataset-items", h.listItems)
router.POST("/dataset-items", h.createItem)
router.PATCH("/dataset-items/:id", h.updateItem)
router.POST("/dataset-items/:id/deposit", h.depositItem)
router.GET("/dataset-crons", h.listCrons)
router.POST("/dataset-crons/sync", h.queueSync)
}
@@ -134,18 +146,21 @@ func (h *Handler) updateSource(c *gin.Context) {
if !ok {
return
}
var input sourceRequest
var input struct {
Name *string `json:"name"`
Kind *string `json:"kind"`
URL *string `json:"url"`
IconURL *string `json:"iconUrl"`
Description *string `json:"description"`
Enabled *bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&input); err != nil {
writeInvalidRequest(c)
return
}
enabled := true
if input.Enabled != nil {
enabled = *input.Enabled
}
source, err := h.service.UpdateSource(userID, identity, SourceInput{
source, err := h.service.PatchSource(userID, identity, SourcePatch{
Name: input.Name, Kind: input.Kind, URL: input.URL, IconURL: input.IconURL,
Description: input.Description, Enabled: enabled,
Description: input.Description, Enabled: input.Enabled,
})
if err != nil {
writeError(c, err)
@@ -180,16 +195,34 @@ func (h *Handler) listItems(c *gin.Context) {
if !ok {
return
}
items, err := h.service.ListItems(userID, c.Query("sourceId"))
limit, err := optionalNonNegativeInt(c.Query("limit"))
if err != nil {
writeInvalidRequest(c)
return
}
offset, err := optionalNonNegativeInt(c.Query("offset"))
if err != nil {
writeInvalidRequest(c)
return
}
page, err := h.service.ListItemsPage(userID, c.Query("sourceId"), limit, offset)
if err != nil {
writeError(c, err)
return
}
result := make([]ItemDTO, 0, len(items))
for _, item := range items {
result := make([]ItemDTO, 0, len(page.Items))
for _, item := range page.Items {
result = append(result, itemDTO(item))
}
c.JSON(http.StatusOK, result)
_, hasLimit := c.GetQuery("limit")
_, hasOffset := c.GetQuery("offset")
if !hasLimit && !hasOffset {
c.JSON(http.StatusOK, result)
return
}
c.JSON(http.StatusOK, ItemPageDTO{
Items: result, Total: page.Total, Limit: page.Limit, Offset: page.Offset,
})
}
func (h *Handler) createItem(c *gin.Context) {
@@ -235,14 +268,14 @@ func (h *Handler) updateItem(c *gin.Context) {
return
}
var input struct {
Status string `json:"status"`
Starred bool `json:"starred"`
Status *string `json:"status"`
Starred *bool `json:"starred"`
}
if err := c.ShouldBindJSON(&input); err != nil {
writeInvalidRequest(c)
return
}
item, err := h.service.UpdateItem(userID, identity, ItemUpdate{Status: input.Status, Starred: input.Starred})
item, err := h.service.PatchItem(userID, identity, ItemPatch{Status: input.Status, Starred: input.Starred})
if err != nil {
writeError(c, err)
return
@@ -250,6 +283,32 @@ func (h *Handler) updateItem(c *gin.Context) {
c.JSON(http.StatusOK, itemDTO(*item))
}
func (h *Handler) depositItem(c *gin.Context) {
userID, ok := currentUser(c)
if !ok {
return
}
identity, ok := httpx.IdentityParam(c, "id")
if !ok {
return
}
var input struct {
ProjectID string `json:"projectId"`
}
if err := c.ShouldBindJSON(&input); err != nil || strings.TrimSpace(input.ProjectID) == "" {
writeInvalidRequest(c)
return
}
result, err := h.service.DepositItem(userID, identity, input.ProjectID)
if err != nil {
writeError(c, err)
return
}
c.JSON(http.StatusCreated, DepositDTO{
NoteID: result.NoteIdentity, ProjectID: result.ProjectIdentity,
})
}
func (h *Handler) listCrons(c *gin.Context) {
userID, ok := currentUser(c)
if !ok {
@@ -272,7 +331,7 @@ func (h *Handler) queueSync(c *gin.Context) {
if !ok {
return
}
crons, err := h.service.SyncSources(c.Request.Context(), userID)
crons, err := h.service.QueueSources(userID)
if err != nil {
writeError(c, err)
return
@@ -281,7 +340,7 @@ func (h *Handler) queueSync(c *gin.Context) {
for _, cron := range crons {
result = append(result, cronDTO(cron))
}
c.JSON(http.StatusOK, result)
c.JSON(http.StatusAccepted, result)
}
func currentUser(c *gin.Context) (uint, bool) {
@@ -297,6 +356,10 @@ func writeError(c *gin.Context, err error) {
switch {
case errors.Is(err, gorm.ErrRecordNotFound):
httpx.Error(c, http.StatusNotFound, "not_found", "数据源或数据条目不存在")
case errors.Is(err, ErrBuiltInSource):
httpx.Error(c, http.StatusConflict, "built_in_source", "内置数据源不能删除,可以将其停用")
case errors.Is(err, ErrSyncInProgress):
httpx.Error(c, http.StatusConflict, "sync_in_progress", "已有数据源采集任务正在运行")
case errors.Is(err, ErrNameRequired), errors.Is(err, ErrKindInvalid),
errors.Is(err, ErrURLInvalid), errors.Is(err, ErrTitleRequired), errors.Is(err, ErrStatusInvalid):
writeInvalidRequest(c)
@@ -313,7 +376,7 @@ func sourceDTO(source models.SaDatasetSource, itemCount int64) SourceDTO {
return SourceDTO{
ID: source.Identity, Name: source.Name, Kind: source.Kind, URL: source.URL,
IconURL: source.IconURL, Description: source.Description, Enabled: source.Enabled,
LastSyncedAt: utcTime(source.LastSyncedAt), ItemCount: itemCount,
BuiltIn: source.SeedKey != nil, LastSyncedAt: utcTime(source.LastSyncedAt), ItemCount: itemCount,
CreatedAt: source.CreatedAt.UTC(), UpdatedAt: source.UpdatedAt.UTC(),
}
}
@@ -329,13 +392,23 @@ func itemDTO(item models.SaDatasetItem) ItemDTO {
func cronDTO(cron models.SaDatasetCron) CronDTO {
return CronDTO{
ID: cron.Identity, SourceID: cron.SourceIdentity, Schedule: cron.Schedule,
Status: cron.Status, Enabled: cron.Enabled, NextRunAt: utcTime(cron.NextRunAt),
ID: cron.Identity, SourceID: cron.SourceIdentity, Status: cron.Status,
LastRunAt: utcTime(cron.LastRunAt), LastResult: cron.LastResult,
CreatedAt: cron.CreatedAt.UTC(), UpdatedAt: cron.UpdatedAt.UTC(),
}
}
func optionalNonNegativeInt(value string) (int, error) {
if strings.TrimSpace(value) == "" {
return 0, nil
}
parsed, err := strconv.Atoi(value)
if err != nil || parsed < 0 {
return 0, errors.New("invalid non-negative integer")
}
return parsed, nil
}
func parseOptionalTime(value string) (*time.Time, error) {
if strings.TrimSpace(value) == "" {
return nil, nil

View File

@@ -8,6 +8,7 @@ import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
@@ -21,12 +22,18 @@ func TestDatasetAPIConnectsSourcesItemsAndSyncTasks(t *testing.T) {
database := newDatasetTestDatabase(t)
owner := createDatasetTestUser(t, database, "owner@example.com")
other := createDatasetTestUser(t, database, "other@example.com")
project := models.SaProject{OwnerID: owner.ID, Name: "Research"}
require.NoError(t, database.Create(&project).Error)
ownerRouter := datasetTestRouter(database, owner.ID)
invalid := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-sources", map[string]any{
"name": "Invalid RSS", "kind": "rss", "url": "",
})
require.Equal(t, http.StatusBadRequest, invalid.Code)
unsupported := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-sources", map[string]any{
"name": "Manual source", "kind": "manual", "enabled": true,
})
require.Equal(t, http.StatusBadRequest, unsupported.Code)
createdSource := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-sources", map[string]any{
"name": "Industry feed", "kind": "rss", "url": "https://example.com/feed",
@@ -53,6 +60,11 @@ func TestDatasetAPIConnectsSourcesItemsAndSyncTasks(t *testing.T) {
var item ItemDTO
require.NoError(t, json.Unmarshal(createdItem.Body.Bytes(), &item))
require.Equal(t, source.ID, item.SourceID)
legacyItems := performDatasetRequest(t, ownerRouter, http.MethodGet, "/api/v1/dataset-items", nil)
require.Equal(t, http.StatusOK, legacyItems.Code)
var legacyItemList []ItemDTO
require.NoError(t, json.Unmarshal(legacyItems.Body.Bytes(), &legacyItemList))
require.Len(t, legacyItemList, 1)
listSources := performDatasetRequest(t, ownerRouter, http.MethodGet, "/api/v1/dataset-sources", nil)
require.Equal(t, http.StatusOK, listSources.Code)
@@ -68,20 +80,42 @@ func TestDatasetAPIConnectsSourcesItemsAndSyncTasks(t *testing.T) {
require.NoError(t, json.Unmarshal(updatedItem.Body.Bytes(), &item))
require.Equal(t, "read", item.Status)
require.True(t, item.Starred)
partialItemUpdate := performDatasetRequest(t, ownerRouter, http.MethodPatch, "/api/v1/dataset-items/"+item.ID, map[string]any{
"starred": false,
})
require.Equal(t, http.StatusOK, partialItemUpdate.Code)
require.NoError(t, json.Unmarshal(partialItemUpdate.Body.Bytes(), &item))
require.Equal(t, "read", item.Status)
require.False(t, item.Starred)
deposited := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-items/"+item.ID+"/deposit", map[string]any{
"projectId": project.Identity,
})
require.Equal(t, http.StatusCreated, deposited.Code)
var depositedNote models.SaNote
require.NoError(t, database.Where("project_id = ?", project.ID).First(&depositedNote).Error)
require.Equal(t, item.Title, depositedNote.Title)
require.Contains(t, depositedNote.Markdown, "原文链接")
firstSync := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-crons/sync", map[string]any{})
require.Equal(t, http.StatusOK, firstSync.Code)
require.Equal(t, http.StatusAccepted, firstSync.Code)
var firstCrons []CronDTO
require.NoError(t, json.Unmarshal(firstSync.Body.Bytes(), &firstCrons))
require.Len(t, firstCrons, 1)
require.Equal(t, "completed", firstCrons[0].Status)
require.Contains(t, firstCrons[0].LastResult, "format=rss")
require.Equal(t, "pending", firstCrons[0].Status)
firstCompleted := waitForDatasetRun(t, ownerRouter, firstCrons[0].ID)
require.Equal(t, "completed", firstCompleted.Status)
require.Contains(t, firstCompleted.LastResult, "format=rss")
secondSync := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-crons/sync", map[string]any{})
require.Equal(t, http.StatusOK, secondSync.Code)
var secondSync *httptest.ResponseRecorder
require.Eventually(t, func() bool {
secondSync = performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-crons/sync", map[string]any{})
return secondSync.Code == http.StatusAccepted
}, time.Second, 10*time.Millisecond)
var secondCrons []CronDTO
require.NoError(t, json.Unmarshal(secondSync.Body.Bytes(), &secondCrons))
require.NotEqual(t, firstCrons[0].ID, secondCrons[0].ID)
waitForDatasetRun(t, ownerRouter, secondCrons[0].ID)
var collectedItems []models.SaDatasetItem
require.NoError(t, database.Where("source_id = ?", 1).Find(&collectedItems).Error)
require.Len(t, collectedItems, 2)
@@ -121,6 +155,84 @@ func TestDatasetSourceAuthorizationUsesOwner(t *testing.T) {
creatorSources, err := NewService(database).ListSources(creator.ID)
require.NoError(t, err)
require.Empty(t, creatorSources)
item := models.SaDatasetItem{
SourceID: source.ID, CreatedBy: creator.ID, Title: "Creator-authored item", Status: "unread",
}
require.NoError(t, database.Create(&item).Error)
page, err := NewService(database).ListItemsPage(owner.ID, source.Identity, 50, 0)
require.NoError(t, err)
require.Len(t, page.Items, 1)
creatorPage, err := NewService(database).ListItemsPage(creator.ID, "", 50, 0)
require.NoError(t, err)
require.Empty(t, creatorPage.Items)
creatorProject := models.SaProject{OwnerID: creator.ID, Name: "Creator project"}
require.NoError(t, database.Create(&creatorProject).Error)
_, err = NewService(database).DepositItem(owner.ID, item.Identity, creatorProject.Identity)
require.ErrorIs(t, err, gorm.ErrRecordNotFound)
require.NoError(t, NewService(database).DeleteSource(owner.ID, source.Identity))
var itemCount int64
require.NoError(t, database.Model(&models.SaDatasetItem{}).Count(&itemCount).Error)
require.Zero(t, itemCount)
}
func TestDatasetItemsArePaginatedByOwnedSource(t *testing.T) {
database := newDatasetTestDatabase(t)
owner := createDatasetTestUser(t, database, "owner@example.com")
source := models.SaDatasetSource{
OwnerID: owner.ID, CreatedBy: owner.ID,
Name: "Feed", Kind: "rss", URL: "https://example.com/feed", Enabled: true,
}
require.NoError(t, database.Create(&source).Error)
for index := 0; index < 3; index++ {
require.NoError(t, database.Create(&models.SaDatasetItem{
SourceID: source.ID, CreatedBy: owner.ID,
Title: fmt.Sprintf("Item %d", index), Status: "unread",
}).Error)
}
router := datasetTestRouter(database, owner.ID)
response := performDatasetRequest(t, router, http.MethodGet,
"/api/v1/dataset-items?sourceId="+source.Identity+"&limit=2&offset=1", nil)
require.Equal(t, http.StatusOK, response.Code)
var page ItemPageDTO
require.NoError(t, json.Unmarshal(response.Body.Bytes(), &page))
require.Equal(t, int64(3), page.Total)
require.Equal(t, 2, page.Limit)
require.Equal(t, 1, page.Offset)
require.Len(t, page.Items, 2)
}
func TestBuiltInDatasetSourceKeepsManagedFieldsAndCannotBeDeleted(t *testing.T) {
database := newDatasetTestDatabase(t)
owner := createDatasetTestUser(t, database, "owner@example.com")
seedKey := "built-in"
source := models.SaDatasetSource{
OwnerID: owner.ID, CreatedBy: owner.ID, SeedKey: &seedKey,
Name: "Built-in", Kind: "rss", URL: "https://example.com/feed",
IconURL: "https://example.com/icon.png", Enabled: true,
}
require.NoError(t, database.Create(&source).Error)
service := NewService(database)
updated, err := service.UpdateSource(owner.ID, source.Identity, SourceInput{
Name: "Renamed", Kind: "rss", URL: "https://attacker.example/feed",
IconURL: "https://attacker.example/icon.png", Enabled: false,
})
require.NoError(t, err)
require.Equal(t, "Renamed", updated.Name)
require.Equal(t, "https://example.com/feed", updated.URL)
require.Equal(t, "https://example.com/icon.png", updated.IconURL)
require.False(t, updated.Enabled)
renamed := "Renamed again"
updated, err = service.PatchSource(owner.ID, source.Identity, SourcePatch{Name: &renamed})
require.NoError(t, err)
require.Equal(t, renamed, updated.Name)
require.False(t, updated.Enabled)
require.Equal(t, "https://example.com/feed", updated.URL)
require.ErrorIs(t, service.DeleteSource(owner.ID, source.Identity), ErrBuiltInSource)
}
func newDatasetTestDatabase(t *testing.T) *gorm.DB {
@@ -132,14 +244,39 @@ func newDatasetTestDatabase(t *testing.T) *gorm.DB {
&models.SaDatasetSource{},
&models.SaDatasetItem{},
&models.SaDatasetCron{},
&models.SaProject{},
&models.SaNote{},
))
require.NoError(t, database.Exec(`
CREATE UNIQUE INDEX idx_sa_dataset_item_source_external
CREATE UNIQUE INDEX IF NOT EXISTS idx_sa_dataset_item_source_external
ON sa_dataset_items (source_id, external_id)
`).Error)
return database
}
func waitForDatasetRun(t *testing.T, router http.Handler, identity string) CronDTO {
t.Helper()
var completed CronDTO
require.Eventually(t, func() bool {
response := performDatasetRequest(t, router, http.MethodGet, "/api/v1/dataset-crons", nil)
if response.Code != http.StatusOK {
return false
}
var runs []CronDTO
if json.Unmarshal(response.Body.Bytes(), &runs) != nil {
return false
}
for _, run := range runs {
if run.ID == identity && (run.Status == "completed" || run.Status == "failed") {
completed = run
return true
}
}
return false
}, time.Second, 10*time.Millisecond)
return completed
}
func createDatasetTestUser(t *testing.T, database *gorm.DB, email string) models.SaUser {
t.Helper()
user := models.SaUser{Email: email, DisplayName: email, PasswordHash: "hash", Role: "user"}

View File

@@ -15,11 +15,20 @@ import (
)
var (
ErrNameRequired = errors.New("dataset source name is required")
ErrKindInvalid = errors.New("dataset source kind is invalid")
ErrURLInvalid = errors.New("dataset source url is invalid")
ErrTitleRequired = errors.New("dataset item title is required")
ErrStatusInvalid = errors.New("dataset item status is invalid")
ErrNameRequired = errors.New("dataset source name is required")
ErrKindInvalid = errors.New("dataset source kind is invalid")
ErrURLInvalid = errors.New("dataset source url is invalid")
ErrTitleRequired = errors.New("dataset item title is required")
ErrStatusInvalid = errors.New("dataset item status is invalid")
ErrBuiltInSource = errors.New("built-in dataset source cannot be deleted")
ErrSyncInProgress = errors.New("dataset sync is already in progress")
)
const (
defaultItemPageSize = 50
maxItemPageSize = 200
datasetRunRetention = 30 * 24 * time.Hour
maxConcurrentFetches = 4
)
type Service struct {
@@ -37,6 +46,15 @@ type SourceInput struct {
Enabled bool
}
type SourcePatch struct {
Name *string
Kind *string
URL *string
IconURL *string
Description *string
Enabled *bool
}
type ItemInput struct {
SourceIdentity string
Title string
@@ -51,11 +69,33 @@ type ItemUpdate struct {
Starred bool
}
type ItemPatch struct {
Status *string
Starred *bool
}
type DepositResult struct {
NoteIdentity string
ProjectIdentity string
}
type SourceRecord struct {
Source models.SaDatasetSource
ItemCount int64
}
type ItemPage struct {
Items []models.SaDatasetItem
Total int64
Limit int
Offset int
}
type queuedSource struct {
Source models.SaDatasetSource
Run models.SaDatasetCron
}
func NewService(database *gorm.DB) *Service {
return &Service{db: database, fetcher: NewHTTPFeedFetcher()}
}
@@ -81,7 +121,7 @@ func (s *Service) ListSources(userID uint) ([]SourceRecord, error) {
}
if err := s.db.Model(&models.SaDatasetItem{}).
Select("source_id, COUNT(*) AS count").
Where("created_by = ? AND source_id IN ?", userID, sourceIDs).
Where("source_id IN ?", sourceIDs).
Group("source_id").
Scan(&rows).Error; err != nil {
return nil, err
@@ -111,14 +151,60 @@ func (s *Service) CreateSource(userID uint, input SourceInput) (*models.SaDatase
}
func (s *Service) UpdateSource(userID uint, identity string, input SourceInput) (*models.SaDatasetSource, error) {
normalized, err := normalizeSourceInput(input)
if err != nil {
return nil, err
return s.PatchSource(userID, identity, SourcePatch{
Name: &input.Name, Kind: &input.Kind, URL: &input.URL, IconURL: &input.IconURL,
Description: &input.Description, Enabled: &input.Enabled,
})
}
func (s *Service) PatchSource(userID uint, identity string, patch SourcePatch) (*models.SaDatasetSource, error) {
if !s.syncMu.TryLock() {
return nil, ErrSyncInProgress
}
defer s.syncMu.Unlock()
var source models.SaDatasetSource
if err := s.db.Where("identity = ? AND owner_id = ?", identity, userID).First(&source).Error; err != nil {
return nil, err
}
input := SourceInput{
Name: source.Name, Kind: source.Kind, URL: source.URL, IconURL: source.IconURL,
Description: source.Description, Enabled: source.Enabled,
}
if patch.Name != nil {
input.Name = *patch.Name
}
if patch.Kind != nil {
input.Kind = *patch.Kind
}
if patch.URL != nil {
input.URL = *patch.URL
}
if patch.IconURL != nil {
input.IconURL = *patch.IconURL
}
if patch.Description != nil {
input.Description = *patch.Description
}
if patch.Enabled != nil {
input.Enabled = *patch.Enabled
}
managedFields := source.SeedKey != nil || source.Kind != "rss"
if managedFields {
input.Kind = source.Kind
input.URL = source.URL
input.IconURL = source.IconURL
}
var normalized SourceInput
var err error
if source.Kind == "rss" {
normalized, err = normalizeSourceInput(input)
} else {
normalized, err = normalizeLegacySourceInput(input)
}
if err != nil {
return nil, err
}
if err := s.db.Model(&source).Updates(map[string]any{
"name": normalized.Name, "kind": normalized.Kind, "url": normalized.URL,
"icon_url": normalized.IconURL, "description": normalized.Description, "enabled": normalized.Enabled,
@@ -129,15 +215,23 @@ func (s *Service) UpdateSource(userID uint, identity string, input SourceInput)
}
func (s *Service) DeleteSource(userID uint, identity string) error {
if !s.syncMu.TryLock() {
return ErrSyncInProgress
}
defer s.syncMu.Unlock()
return s.db.Transaction(func(tx *gorm.DB) error {
var source models.SaDatasetSource
if err := tx.Where("identity = ? AND owner_id = ?", identity, userID).First(&source).Error; err != nil {
return err
}
if err := tx.Where("source_id = ? AND created_by = ?", source.ID, userID).Delete(&models.SaDatasetCron{}).Error; err != nil {
if source.SeedKey != nil {
return ErrBuiltInSource
}
if err := tx.Where("source_id = ?", source.ID).Delete(&models.SaDatasetCron{}).Error; err != nil {
return err
}
if err := tx.Where("source_id = ? AND created_by = ?", source.ID, userID).Delete(&models.SaDatasetItem{}).Error; err != nil {
if err := tx.Where("source_id = ?", source.ID).Delete(&models.SaDatasetItem{}).Error; err != nil {
return err
}
return tx.Delete(&source).Error
@@ -145,28 +239,46 @@ func (s *Service) DeleteSource(userID uint, identity string) error {
}
func (s *Service) ListItems(userID uint, sourceIdentity string) ([]models.SaDatasetItem, error) {
query := s.db.Where("created_by = ?", userID)
page, err := s.ListItemsPage(userID, sourceIdentity, maxItemPageSize, 0)
if err != nil {
return nil, err
}
return page.Items, nil
}
func (s *Service) ListItemsPage(userID uint, sourceIdentity string, limit, offset int) (ItemPage, error) {
limit, offset = normalizePage(limit, offset)
query := s.db.Model(&models.SaDatasetItem{}).
Joins("JOIN sa_dataset_sources ON sa_dataset_sources.id = sa_dataset_items.source_id").
Where("sa_dataset_sources.owner_id = ?", userID)
if sourceIdentity = strings.TrimSpace(sourceIdentity); sourceIdentity != "" {
source, err := s.findOwnedSource(userID, sourceIdentity)
if err != nil {
return nil, err
return ItemPage{}, err
}
query = query.Where("source_id = ?", source.ID)
query = query.Where("sa_dataset_items.source_id = ?", source.ID)
}
var total int64
if err := query.Session(&gorm.Session{}).Count(&total).Error; err != nil {
return ItemPage{}, err
}
var items []models.SaDatasetItem
if err := query.Order("COALESCE(published_at, created_at) desc, id desc").Limit(200).Find(&items).Error; err != nil {
return nil, err
if err := query.Select("sa_dataset_items.*").
Order("COALESCE(sa_dataset_items.published_at, sa_dataset_items.created_at) desc, sa_dataset_items.id desc").
Limit(limit).Offset(offset).Find(&items).Error; err != nil {
return ItemPage{}, err
}
if items == nil {
items = []models.SaDatasetItem{}
}
return items, nil
return ItemPage{Items: items, Total: total, Limit: limit, Offset: offset}, nil
}
func (s *Service) CountItems(userID uint, sourceID uint) (int64, error) {
var count int64
err := s.db.Model(&models.SaDatasetItem{}).
Where("created_by = ? AND source_id = ?", userID, sourceID).
Joins("JOIN sa_dataset_sources ON sa_dataset_sources.id = sa_dataset_items.source_id").
Where("sa_dataset_sources.owner_id = ? AND sa_dataset_items.source_id = ?", userID, sourceID).
Count(&count).Error
return count, err
}
@@ -193,23 +305,72 @@ func (s *Service) CreateItem(userID uint, input ItemInput) (*models.SaDatasetIte
}
func (s *Service) UpdateItem(userID uint, identity string, input ItemUpdate) (*models.SaDatasetItem, error) {
status := strings.TrimSpace(input.Status)
if status != "unread" && status != "read" && status != "archived" {
return nil, ErrStatusInvalid
}
var item models.SaDatasetItem
if err := s.db.Where("identity = ? AND created_by = ?", identity, userID).First(&item).Error; err != nil {
return s.PatchItem(userID, identity, ItemPatch{Status: &input.Status, Starred: &input.Starred})
}
func (s *Service) PatchItem(userID uint, identity string, patch ItemPatch) (*models.SaDatasetItem, error) {
item, err := s.findOwnedItem(userID, identity)
if err != nil {
return nil, err
}
if err := s.db.Model(&item).Updates(map[string]any{"status": status, "starred": input.Starred}).Error; err != nil {
updates := make(map[string]any, 2)
if patch.Status != nil {
status := strings.TrimSpace(*patch.Status)
if status != "unread" && status != "read" && status != "archived" {
return nil, ErrStatusInvalid
}
updates["status"] = status
}
if patch.Starred != nil {
updates["starred"] = *patch.Starred
}
if len(updates) == 0 {
return item, nil
}
if err := s.db.Model(item).Updates(updates).Error; err != nil {
return nil, err
}
return &item, nil
return item, nil
}
func (s *Service) DepositItem(userID uint, itemIdentity, projectIdentity string) (DepositResult, error) {
item, err := s.findOwnedItem(userID, itemIdentity)
if err != nil {
return DepositResult{}, err
}
var project models.SaProject
if err := s.db.Where("identity = ? AND owner_id = ?", strings.TrimSpace(projectIdentity), userID).
First(&project).Error; err != nil {
return DepositResult{}, err
}
body := strings.TrimSpace(item.Content)
if body == "" {
body = strings.TrimSpace(item.Summary)
}
if item.URL != "" {
if body != "" {
body += "\n\n"
}
body += "[原文链接](" + item.URL + ")"
}
note := models.SaNote{
ProjectID: project.ID, CreatedBy: userID,
Title: item.Title, Markdown: body,
}
if err := s.db.Create(&note).Error; err != nil {
return DepositResult{}, err
}
return DepositResult{NoteIdentity: note.Identity, ProjectIdentity: project.Identity}, nil
}
func (s *Service) ListCrons(userID uint) ([]models.SaDatasetCron, error) {
var crons []models.SaDatasetCron
if err := s.db.Where("created_by = ?", userID).Order("created_at desc, id desc").Limit(100).Find(&crons).Error; err != nil {
if err := s.db.Model(&models.SaDatasetCron{}).
Select("sa_dataset_crons.*").
Joins("JOIN sa_dataset_sources ON sa_dataset_sources.id = sa_dataset_crons.source_id").
Where("sa_dataset_sources.owner_id = ?", userID).
Order("sa_dataset_crons.created_at desc, sa_dataset_crons.id desc").
Limit(100).Find(&crons).Error; err != nil {
return nil, err
}
if crons == nil {
@@ -222,7 +383,36 @@ func (s *Service) SyncSources(ctx context.Context, userID uint) ([]models.SaData
s.syncMu.Lock()
defer s.syncMu.Unlock()
return s.syncSources(ctx, userID)
return s.syncSources(ctx, userID, "scheduled")
}
func (s *Service) QueueSources(userID uint) ([]models.SaDatasetCron, error) {
if !s.syncMu.TryLock() {
return nil, ErrSyncInProgress
}
sources, err := s.enabledRSSSources(userID)
if err != nil {
s.syncMu.Unlock()
return nil, err
}
queued, err := s.createDatasetRuns(userID, sources, "manual")
if err != nil {
s.syncMu.Unlock()
return nil, err
}
runs := make([]models.SaDatasetCron, 0, len(queued))
for _, entry := range queued {
runs = append(runs, entry.Run)
}
if len(queued) == 0 {
s.syncMu.Unlock()
return runs, nil
}
go func() {
defer s.syncMu.Unlock()
_, _ = s.executeDatasetRuns(context.Background(), queued)
}()
return runs, nil
}
func (s *Service) SyncAllSources(ctx context.Context) ([]models.SaDatasetCron, error) {
@@ -250,103 +440,166 @@ func (s *Service) SyncAllSources(ctx context.Context) ([]models.SaDatasetCron, e
return result, errors.Join(syncErrors...)
}
func (s *Service) syncSources(ctx context.Context, userID uint) ([]models.SaDatasetCron, error) {
result := make([]models.SaDatasetCron, 0)
func (s *Service) syncSources(ctx context.Context, userID uint, trigger string) ([]models.SaDatasetCron, error) {
sources, err := s.enabledRSSSources(userID)
if err != nil {
return nil, err
}
queued, err := s.createDatasetRuns(userID, sources, trigger)
if err != nil {
return nil, err
}
return s.executeDatasetRuns(ctx, queued)
}
func (s *Service) enabledRSSSources(userID uint) ([]models.SaDatasetSource, error) {
var sources []models.SaDatasetSource
if err := s.db.Where("owner_id = ? AND enabled = ? AND kind = ?", userID, true, "rss").
Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
for _, source := range sources {
if err := ctx.Err(); err != nil {
return result, err
}
now := time.Now().UTC()
cron := models.SaDatasetCron{
SourceID: source.ID, CreatedBy: userID, Schedule: "@once",
Status: "running", Enabled: true, NextRunAt: &now, LastRunAt: &now,
}
if err := s.db.Create(&cron).Error; err != nil {
return nil, err
}
return sources, nil
}
feed, fetchErr := s.fetcher.Fetch(ctx, source.URL)
if fetchErr != nil {
cron.Status = "failed"
cron.Enabled = false
cron.NextRunAt = nil
cron.LastResult = truncateResult(fetchErr.Error())
if err := s.db.Model(&cron).Updates(map[string]any{
"status": cron.Status, "last_result": cron.LastResult, "enabled": false, "next_run_at": nil,
}).Error; err != nil {
return nil, err
}
result = append(result, cron)
if err := ctx.Err(); err != nil {
return result, err
}
continue
func (s *Service) createDatasetRuns(userID uint, sources []models.SaDatasetSource, trigger string) ([]queuedSource, error) {
queued := make([]queuedSource, 0, len(sources))
now := time.Now().UTC()
err := s.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("created_at < ?", now.Add(-datasetRunRetention)).
Delete(&models.SaDatasetCron{}).Error; err != nil {
return err
}
for _, source := range sources {
run := models.SaDatasetCron{
SourceID: source.ID, CreatedBy: userID, Schedule: trigger,
Status: "pending", Enabled: true, NextRunAt: &now,
}
if err := tx.Create(&run).Error; err != nil {
return err
}
queued = append(queued, queuedSource{Source: source, Run: run})
}
return nil
})
return queued, err
}
inserted, err := s.storeFeedItems(userID, source, feed.Items)
if err != nil {
cron.Status = "failed"
cron.Enabled = false
cron.NextRunAt = nil
cron.LastResult = truncateResult("store feed items: " + err.Error())
if updateErr := s.db.Model(&cron).Updates(map[string]any{
"status": cron.Status, "last_result": cron.LastResult, "enabled": false, "next_run_at": nil,
}).Error; updateErr != nil {
return nil, updateErr
}
result = append(result, cron)
continue
func (s *Service) executeDatasetRuns(ctx context.Context, queued []queuedSource) ([]models.SaDatasetCron, error) {
result := make([]models.SaDatasetCron, len(queued))
runErrors := make([]error, len(queued))
concurrency := maxConcurrentFetches
if s.db.Dialector.Name() == "sqlite" {
concurrency = 1
}
semaphore := make(chan struct{}, concurrency)
var workers sync.WaitGroup
for index, entry := range queued {
workers.Add(1)
go func() {
defer workers.Done()
semaphore <- struct{}{}
defer func() { <-semaphore }()
result[index], runErrors[index] = s.executeDatasetRun(ctx, entry)
}()
}
workers.Wait()
return result, errors.Join(runErrors...)
}
func (s *Service) executeDatasetRun(ctx context.Context, entry queuedSource) (models.SaDatasetCron, error) {
source := entry.Source
cron := entry.Run
if err := ctx.Err(); err != nil {
return cron, err
}
now := time.Now().UTC()
cron.Status = "running"
cron.LastRunAt = &now
if err := s.db.Model(&cron).Updates(map[string]any{
"status": "running", "last_run_at": now,
}).Error; err != nil {
return cron, err
}
feed, fetchErr := s.fetcher.Fetch(ctx, source.URL)
if fetchErr != nil {
cron.Status = "failed"
cron.Enabled = false
cron.NextRunAt = nil
cron.LastResult = truncateResult(fetchErr.Error())
if err := s.db.Model(&cron).Updates(map[string]any{
"status": cron.Status, "last_result": cron.LastResult, "enabled": false, "next_run_at": nil,
}).Error; err != nil {
return cron, err
}
return cron, ctx.Err()
}
inserted := 0
completedAt := time.Now().UTC()
err := s.db.Transaction(func(tx *gorm.DB) error {
var storeErr error
inserted, storeErr = storeFeedItems(tx, cron.CreatedBy, source, feed.Items)
if storeErr != nil {
return storeErr
}
completedAt := time.Now().UTC()
cron.Status = "completed"
cron.Enabled = false
cron.NextRunAt = nil
cron.LastResult = fmt.Sprintf("format=%s fetched=%d inserted=%d", feed.Format, len(feed.Items), inserted)
if err := s.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&cron).Updates(map[string]any{
"status": cron.Status, "last_result": cron.LastResult, "enabled": false, "next_run_at": nil,
}).Error; err != nil {
return err
}
return tx.Model(&source).Update("last_synced_at", completedAt).Error
}); err != nil {
return nil, err
if err := tx.Model(&cron).Updates(map[string]any{
"status": cron.Status, "last_result": cron.LastResult, "enabled": false, "next_run_at": nil,
}).Error; err != nil {
return err
}
source.LastSyncedAt = &completedAt
result = append(result, cron)
return tx.Model(&source).Update("last_synced_at", completedAt).Error
})
if err == nil {
return cron, nil
}
return result, nil
cron.Status = "failed"
cron.Enabled = false
cron.NextRunAt = nil
cron.LastResult = truncateResult("store feed items: " + err.Error())
if updateErr := s.db.Model(&cron).Updates(map[string]any{
"status": cron.Status, "last_result": cron.LastResult, "enabled": false, "next_run_at": nil,
}).Error; updateErr != nil {
return cron, updateErr
}
return cron, nil
}
func (s *Service) storeFeedItems(userID uint, source models.SaDatasetSource, items []FeedItem) (int, error) {
inserted := 0
err := s.db.Transaction(func(tx *gorm.DB) error {
for _, input := range items {
externalID := input.ExternalID
item := models.SaDatasetItem{
SourceID: source.ID, CreatedBy: userID, ExternalID: &externalID,
Title: input.Title, Summary: input.Summary, Content: input.Content,
URL: input.URL, Status: "unread", PublishedAt: input.PublishedAt,
}
result := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "source_id"}, {Name: "external_id"}},
DoNothing: true,
}).Create(&item)
if result.Error != nil {
return result.Error
}
inserted += int(result.RowsAffected)
}
return nil
var err error
inserted, err = storeFeedItems(tx, userID, source, items)
return err
})
return inserted, err
}
func storeFeedItems(tx *gorm.DB, userID uint, source models.SaDatasetSource, items []FeedItem) (int, error) {
inserted := 0
for _, input := range items {
externalID := input.ExternalID
item := models.SaDatasetItem{
SourceID: source.ID, CreatedBy: userID, ExternalID: &externalID,
Title: input.Title, Summary: input.Summary, Content: input.Content,
URL: input.URL, Status: "unread", PublishedAt: input.PublishedAt,
}
result := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "source_id"}, {Name: "external_id"}},
DoNothing: true,
}).Create(&item)
if result.Error != nil {
return inserted, result.Error
}
inserted += int(result.RowsAffected)
}
return inserted, nil
}
func truncateResult(value string) string {
return truncateRunes(value, 2000)
}
@@ -365,6 +618,16 @@ func (s *Service) findOwnedSource(userID uint, identity string) (*models.SaDatas
return &source, err
}
func (s *Service) findOwnedItem(userID uint, identity string) (*models.SaDatasetItem, error) {
var item models.SaDatasetItem
err := s.db.Model(&models.SaDatasetItem{}).
Select("sa_dataset_items.*").
Joins("JOIN sa_dataset_sources ON sa_dataset_sources.id = sa_dataset_items.source_id").
Where("sa_dataset_items.identity = ? AND sa_dataset_sources.owner_id = ?", strings.TrimSpace(identity), userID).
First(&item).Error
return &item, err
}
func normalizeSourceInput(input SourceInput) (SourceInput, error) {
input.Name = strings.TrimSpace(input.Name)
input.Kind = strings.ToLower(strings.TrimSpace(input.Kind))
@@ -374,7 +637,7 @@ func normalizeSourceInput(input SourceInput) (SourceInput, error) {
if input.Name == "" {
return input, ErrNameRequired
}
if input.Kind != "manual" && input.Kind != "link" && input.Kind != "rss" {
if input.Kind != "rss" {
return input, ErrKindInvalid
}
if input.URL != "" && !validHTTPURL(input.URL) {
@@ -383,12 +646,34 @@ func normalizeSourceInput(input SourceInput) (SourceInput, error) {
if input.IconURL != "" && !validHTTPURL(input.IconURL) {
return input, ErrURLInvalid
}
if input.Kind != "manual" && input.URL == "" {
if input.URL == "" {
return input, ErrURLInvalid
}
return input, nil
}
func normalizeLegacySourceInput(input SourceInput) (SourceInput, error) {
input.Name = strings.TrimSpace(input.Name)
input.Description = strings.TrimSpace(input.Description)
if input.Name == "" {
return input, ErrNameRequired
}
return input, nil
}
func normalizePage(limit, offset int) (int, int) {
if limit <= 0 {
limit = defaultItemPageSize
}
if limit > maxItemPageSize {
limit = maxItemPageSize
}
if offset < 0 {
offset = 0
}
return limit, offset
}
func validHTTPURL(value string) bool {
parsed, err := url.ParseRequestURI(value)
return err == nil && (parsed.Scheme == "http" || parsed.Scheme == "https") &&