327 lines
13 KiB
Go
327 lines
13 KiB
Go
package dataset
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/gorm"
|
|
"senlinai-agent/backend/internal/config"
|
|
"senlinai-agent/backend/internal/httpx"
|
|
"senlinai-agent/backend/internal/models"
|
|
)
|
|
|
|
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",
|
|
"description": "Industry updates", "enabled": true,
|
|
})
|
|
require.Equal(t, http.StatusCreated, createdSource.Code)
|
|
var source SourceDTO
|
|
require.NoError(t, json.Unmarshal(createdSource.Body.Bytes(), &source))
|
|
require.NotEmpty(t, source.ID)
|
|
require.Equal(t, "Industry feed", source.Name)
|
|
require.Empty(t, source.IconURL)
|
|
var storedSource models.SaDatasetSource
|
|
require.NoError(t, database.Where("identity = ?", source.ID).First(&storedSource).Error)
|
|
require.Equal(t, owner.ID, storedSource.OwnerID)
|
|
require.Equal(t, owner.Identity, storedSource.OwnerIdentity)
|
|
require.Equal(t, owner.ID, storedSource.CreatedBy)
|
|
require.Equal(t, owner.Identity, storedSource.CreatedByIdentity)
|
|
|
|
createdItem := performDatasetRequest(t, ownerRouter, http.MethodPost, "/api/v1/dataset-items", map[string]any{
|
|
"sourceId": source.ID, "title": "A collected article", "summary": "Summary",
|
|
"url": "https://example.com/article",
|
|
})
|
|
require.Equal(t, http.StatusCreated, createdItem.Code)
|
|
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)
|
|
var sources []SourceDTO
|
|
require.NoError(t, json.Unmarshal(listSources.Body.Bytes(), &sources))
|
|
require.Len(t, sources, 1)
|
|
require.Equal(t, int64(1), sources[0].ItemCount)
|
|
|
|
updatedItem := performDatasetRequest(t, ownerRouter, http.MethodPatch, "/api/v1/dataset-items/"+item.ID, map[string]any{
|
|
"status": "read", "starred": true,
|
|
})
|
|
require.Equal(t, http.StatusOK, updatedItem.Code)
|
|
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.StatusAccepted, firstSync.Code)
|
|
var firstCrons []CronDTO
|
|
require.NoError(t, json.Unmarshal(firstSync.Body.Bytes(), &firstCrons))
|
|
require.Len(t, firstCrons, 1)
|
|
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")
|
|
|
|
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)
|
|
|
|
otherRouter := datasetTestRouter(database, other.ID)
|
|
otherSources := performDatasetRequest(t, otherRouter, http.MethodGet, "/api/v1/dataset-sources", nil)
|
|
require.Equal(t, http.StatusOK, otherSources.Code)
|
|
require.JSONEq(t, `[]`, otherSources.Body.String())
|
|
otherUpdate := performDatasetRequest(t, otherRouter, http.MethodPatch, "/api/v1/dataset-sources/"+source.ID, map[string]any{
|
|
"name": "Stolen", "kind": "manual", "enabled": true,
|
|
})
|
|
require.Equal(t, http.StatusNotFound, otherUpdate.Code)
|
|
|
|
deleted := performDatasetRequest(t, ownerRouter, http.MethodDelete, "/api/v1/dataset-sources/"+source.ID, nil)
|
|
require.Equal(t, http.StatusNoContent, deleted.Code)
|
|
var itemCount, cronCount int64
|
|
require.NoError(t, database.Model(&models.SaDatasetItem{}).Count(&itemCount).Error)
|
|
require.NoError(t, database.Model(&models.SaDatasetCron{}).Count(&cronCount).Error)
|
|
require.Zero(t, itemCount)
|
|
require.Zero(t, cronCount)
|
|
}
|
|
|
|
func TestDatasetSourceAuthorizationUsesOwner(t *testing.T) {
|
|
database := newDatasetTestDatabase(t)
|
|
owner := createDatasetTestUser(t, database, "owner@example.com")
|
|
creator := createDatasetTestUser(t, database, "creator@example.com")
|
|
source := models.SaDatasetSource{
|
|
OwnerID: owner.ID, CreatedBy: creator.ID,
|
|
Name: "Owned feed", Kind: "rss", URL: "https://example.com/feed", Enabled: true,
|
|
}
|
|
require.NoError(t, database.Create(&source).Error)
|
|
|
|
ownerSources, err := NewService(database).ListSources(owner.ID)
|
|
require.NoError(t, err)
|
|
require.Len(t, ownerSources, 1)
|
|
|
|
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 {
|
|
t.Helper()
|
|
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{},
|
|
&models.SaDatasetCron{},
|
|
&models.SaProject{},
|
|
&models.SaNote{},
|
|
))
|
|
require.NoError(t, database.Exec(`
|
|
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"}
|
|
require.NoError(t, database.Create(&user).Error)
|
|
return user
|
|
}
|
|
|
|
func datasetTestRouter(database *gorm.DB, userID uint) http.Handler {
|
|
return httpx.NewProtectedRouter(
|
|
config.Config{Env: "test"},
|
|
func(string) (uint, error) { return userID, nil },
|
|
NewHandler(newServiceWithFetcher(database, staticFeedFetcher{})),
|
|
)
|
|
}
|
|
|
|
type staticFeedFetcher struct{}
|
|
|
|
func (staticFeedFetcher) Fetch(context.Context, string) (ParsedFeed, error) {
|
|
return ParsedFeed{
|
|
Format: "rss",
|
|
Items: []FeedItem{{
|
|
ExternalID: "stable-feed-item",
|
|
Title: "Collected item",
|
|
Summary: "Collected summary",
|
|
Content: "Collected content",
|
|
URL: "https://example.com/collected",
|
|
}},
|
|
}, nil
|
|
}
|
|
|
|
func performDatasetRequest(t *testing.T, router http.Handler, method, path string, body any) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
var encoded []byte
|
|
if body != nil {
|
|
var err error
|
|
encoded, err = json.Marshal(body)
|
|
require.NoError(t, err)
|
|
}
|
|
request := httptest.NewRequest(method, path, bytes.NewReader(encoded))
|
|
request.Header.Set("Authorization", "Bearer test-token")
|
|
if body != nil {
|
|
request.Header.Set("Content-Type", "application/json")
|
|
}
|
|
recorder := httptest.NewRecorder()
|
|
router.ServeHTTP(recorder, request)
|
|
return recorder
|
|
}
|