diff --git a/internal/cli/commands.go b/internal/cli/commands.go index 305e8e1..dc29e38 100644 --- a/internal/cli/commands.go +++ b/internal/cli/commands.go @@ -2,6 +2,7 @@ package cli import ( "bufio" + "context" "fmt" "net/http" "os" @@ -234,14 +235,43 @@ func newArticlesCommand() *cobra.Command { return markError(err) } + limit := viper.GetInt("limit") + if limit < 0 { + err := fmt.Errorf("invalid --limit: %d (must be >= 0)", limit) + printError(err) + return markError(err) + } + if limit > storage.MaxListLimit { + err := fmt.Errorf("invalid --limit: %d (maximum is %d)", limit, storage.MaxListLimit) + printError(err) + return markError(err) + } + return withDatabase(cmd, func(db *storage.Database) error { - articles, blogNames, err := controller.GetArticles(cmd.Context(), db, showAll, viper.GetString("blog"), viper.GetString("category"), since, before) + filter := storage.ArticleFilter{ + UnreadOnly: !showAll, + Category: stringPtr(viper.GetString("category")), + Since: since, + Before: before, + Search: viper.GetString("search"), + Limit: limit, + } + + blogID, err := resolveBlogID(cmd.Context(), db, viper.GetString("blog")) + if err != nil { + return err + } + filter.BlogID = blogID + + articles, blogNames, err := controller.GetArticles(cmd.Context(), db, filter) if err != nil { printError(err) return markError(err) } if len(articles) == 0 { - if showAll { + if viper.GetString("search") != "" { + cprintf([]color.Attribute{color.FgYellow}, "No articles matching '%s'.\n", viper.GetString("search")) + } else if showAll { fmt.Println("No articles found.") } else { cprintln([]color.Attribute{color.FgGreen}, "No unread articles!") @@ -253,6 +283,9 @@ func newArticlesCommand() *cobra.Command { if showAll { label = "All articles" } + if viper.GetString("search") != "" { + label = fmt.Sprintf("Search results for '%s'", viper.GetString("search")) + } cprintf([]color.Attribute{color.FgCyan, color.Bold}, "%s (%d):\n\n", label, len(articles)) for _, article := range articles { printArticle(article, blogNames[article.BlogID]) @@ -267,6 +300,8 @@ func newArticlesCommand() *cobra.Command { cmd.Flags().StringP("category", "c", "", "Filter by category") cmd.Flags().String("since", "", "Show articles published on or after YYYY-MM-DD") cmd.Flags().String("before", "", "Show articles published before YYYY-MM-DD") + cmd.Flags().StringP("search", "s", "", "Search articles by title or content (FTS5 full-text search)") + cmd.Flags().IntP("limit", "n", 20, "Maximum number of articles to return") return cmd } @@ -303,10 +338,17 @@ func newReadAllCommand() *cobra.Command { Use: "read-all", Short: "Mark all unread articles as read.", RunE: func(cmd *cobra.Command, args []string) error { - blogName := viper.GetString("blog") - return withDatabase(cmd, func(db *storage.Database) error { - articles, _, err := controller.GetArticles(cmd.Context(), db, false, blogName, "", nil, nil) + filter := storage.ArticleFilter{ + UnreadOnly: true, + } + blogID, err := resolveBlogID(cmd.Context(), db, viper.GetString("blog")) + if err != nil { + return err + } + filter.BlogID = blogID + + articles, _, err := controller.GetArticles(cmd.Context(), db, filter) if err != nil { printError(err) return markError(err) @@ -318,8 +360,8 @@ func newReadAllCommand() *cobra.Command { if !viper.GetBool("yes") { scope := "all blogs" - if blogName != "" { - scope = fmt.Sprintf("from '%s'", blogName) + if blog := viper.GetString("blog"); blog != "" { + scope = fmt.Sprintf("from '%s'", blog) } confirmed, err := confirm(fmt.Sprintf("Mark %d article(s) %s as read?", len(articles), scope)) if err != nil { @@ -330,7 +372,7 @@ func newReadAllCommand() *cobra.Command { } } - marked, err := controller.MarkAllArticlesRead(cmd.Context(), db, blogName) + marked, err := controller.MarkAllArticlesRead(cmd.Context(), db, filter) if err != nil { printError(err) return markError(err) @@ -508,6 +550,29 @@ func parseDateRange(sinceStr, beforeStr string) (*time.Time, *time.Time, error) return since, before, nil } +func stringPtr(s string) *string { + if s == "" { + return nil + } + return &s +} + +func resolveBlogID(ctx context.Context, db *storage.Database, blogName string) (*int64, error) { + if blogName == "" { + return nil, nil + } + blog, err := db.GetBlogByName(ctx, blogName) + if err != nil { + return nil, err + } + if blog == nil { + err := fmt.Errorf("blog '%s' not found", blogName) + printError(err) + return nil, markError(err) + } + return &blog.ID, nil +} + func confirm(prompt string) (bool, error) { reader := bufio.NewReader(os.Stdin) fmt.Printf("%s [y/N]: ", prompt) diff --git a/internal/controller/controller.go b/internal/controller/controller.go index a6fd8f0..d097b96 100644 --- a/internal/controller/controller.go +++ b/internal/controller/controller.go @@ -7,7 +7,6 @@ import ( "io" "net/url" "strings" - "time" "github.com/JulienTant/blogwatcher-cli/internal/model" "github.com/JulienTant/blogwatcher-cli/internal/opml" @@ -100,38 +99,33 @@ func RemoveBlog(ctx context.Context, db *storage.Database, name string) error { return err } -func GetArticles(ctx context.Context, db *storage.Database, showAll bool, blogName string, category string, since *time.Time, before *time.Time) ([]model.Article, map[int64]string, error) { - var blogID *int64 - if blogName != "" { - blog, err := db.GetBlogByName(ctx, blogName) - if err != nil { - return nil, nil, err - } - if blog == nil { - return nil, nil, BlogNotFoundError{Name: blogName} - } - blogID = &blog.ID - } - - var categoryPtr *string - if category != "" { - categoryPtr = &category +func GetArticles(ctx context.Context, db *storage.Database, filter storage.ArticleFilter) ([]model.Article, map[int64]string, error) { + articles, err := db.ListArticles(ctx, filter) + if err != nil { + return nil, nil, err } - articles, err := db.ListArticles(ctx, !showAll, blogID, categoryPtr, since, before) + blogNames, err := getBlogNames(ctx, db, articles) if err != nil { return nil, nil, err } + + return articles, blogNames, nil +} + +func getBlogNames(ctx context.Context, db *storage.Database, articles []model.Article) (map[int64]string, error) { + if len(articles) == 0 { + return nil, nil + } blogs, err := db.ListBlogs(ctx) if err != nil { - return nil, nil, err + return nil, err } blogNames := make(map[int64]string) for _, blog := range blogs { blogNames[blog.ID] = blog.Name } - - return articles, blogNames, nil + return blogNames, nil } func MarkArticleRead(ctx context.Context, db *storage.Database, articleID int64) (model.Article, error) { @@ -151,29 +145,24 @@ func MarkArticleRead(ctx context.Context, db *storage.Database, articleID int64) return *article, nil } -func MarkAllArticlesRead(ctx context.Context, db *storage.Database, blogName string) ([]model.Article, error) { - var blogID *int64 - if blogName != "" { - blog, err := db.GetBlogByName(ctx, blogName) - if err != nil { - return nil, err - } - if blog == nil { - return nil, BlogNotFoundError{Name: blogName} - } - blogID = &blog.ID - } +func MarkAllArticlesRead(ctx context.Context, db *storage.Database, filter storage.ArticleFilter) ([]model.Article, error) { + filter.UnreadOnly = true - articles, err := db.ListArticles(ctx, true, blogID, nil, nil, nil) + articles, err := db.ListArticles(ctx, filter) if err != nil { return nil, err } + if len(articles) == 0 { + return articles, nil + } - for _, article := range articles { - _, err := db.MarkArticleRead(ctx, article.ID) - if err != nil { - return nil, err - } + ids := make([]int64, len(articles)) + for i, a := range articles { + ids[i] = a.ID + } + + if err := db.MarkArticlesRead(ctx, ids); err != nil { + return nil, err } return articles, nil diff --git a/internal/controller/controller_test.go b/internal/controller/controller_test.go index a90f06d..a914e09 100644 --- a/internal/controller/controller_test.go +++ b/internal/controller/controller_test.go @@ -122,13 +122,10 @@ func TestGetArticlesFilters(t *testing.T) { _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "Title", URL: "https://example.com/1"}) require.NoError(t, err, "add article") - articles, blogNames, err := GetArticles(ctx, db, false, "", "", nil, nil) + articles, blogNames, err := GetArticles(ctx, db, storage.ArticleFilter{}) require.NoError(t, err, "get articles") require.Len(t, articles, 1) require.Equal(t, blog.Name, blogNames[blog.ID]) - - _, _, err = GetArticles(ctx, db, false, "Missing", "", nil, nil) - require.Error(t, err, "expected blog not found error") } func TestImportOPML(t *testing.T) { @@ -290,13 +287,14 @@ func TestGetArticlesFilterByCategory(t *testing.T) { require.NoError(t, err, "add article") // Filter by Go - articles, _, err := GetArticles(ctx, db, false, "", "Go", nil, nil) + cat := "Go" + articles, _, err := GetArticles(ctx, db, storage.ArticleFilter{Category: &cat}) require.NoError(t, err, "get articles by category") require.Len(t, articles, 1) require.Equal(t, "Go Post", articles[0].Title) // No filter returns all - all, _, err := GetArticles(ctx, db, false, "", "", nil, nil) + all, _, err := GetArticles(ctx, db, storage.ArticleFilter{}) require.NoError(t, err, "get all articles") require.Len(t, all, 2) } diff --git a/internal/model/model.go b/internal/model/model.go index 07d0fee..f857611 100644 --- a/internal/model/model.go +++ b/internal/model/model.go @@ -20,4 +20,6 @@ type Article struct { DiscoveredDate *time.Time IsRead bool Categories []string + Description string + Content string } diff --git a/internal/rss/rss.go b/internal/rss/rss.go index df8e650..e401784 100644 --- a/internal/rss/rss.go +++ b/internal/rss/rss.go @@ -20,6 +20,8 @@ type FeedArticle struct { URL string PublishedDate *time.Time Categories []string + Description string + Content string } type FeedParseError struct { @@ -76,6 +78,8 @@ func (f *Fetcher) ParseFeed(ctx context.Context, feedURL string) ([]FeedArticle, URL: link, PublishedDate: pickPublishedDate(item), Categories: item.Categories, + Description: item.Description, + Content: item.Content, }) } diff --git a/internal/scanner/scanner.go b/internal/scanner/scanner.go index 3d8b05e..4257ceb 100644 --- a/internal/scanner/scanner.go +++ b/internal/scanner/scanner.go @@ -4,7 +4,6 @@ import ( "context" "errors" "fmt" - "os" "time" "github.com/cenkalti/backoff/v5" @@ -209,17 +208,8 @@ func (s *Scanner) ScanAllBlogs(ctx context.Context, db *storage.Database, worker for i := 0; i < workers; i++ { g.Go(func() error { - workerDB, openErr := storage.OpenDatabase(gctx, db.Path()) - if openErr != nil { - return openErr - } - defer func() { - if closeErr := workerDB.Close(); closeErr != nil { - fmt.Fprintf(os.Stderr, "close: %v\n", closeErr) - } - }() for item := range jobs { - result, scanErr := s.ScanBlog(gctx, workerDB, item.Blog) + result, scanErr := s.ScanBlog(gctx, db, item.Blog) if scanErr != nil { if isFatalScanError(scanErr) { return fmt.Errorf("scan %s: %w", item.Blog.Name, scanErr) @@ -277,6 +267,8 @@ func convertFeedArticles(blogID int64, articles []rss.FeedArticle) []model.Artic PublishedDate: article.PublishedDate, IsRead: false, Categories: article.Categories, + Description: article.Description, + Content: article.Content, }) } return result diff --git a/internal/scanner/scanner_test.go b/internal/scanner/scanner_test.go index 92ce27f..24c27d1 100644 --- a/internal/scanner/scanner_test.go +++ b/internal/scanner/scanner_test.go @@ -58,7 +58,7 @@ func TestScanBlogRSS(t *testing.T) { require.Equal(t, 2, result.NewArticles) require.Equal(t, "rss", result.Source) - articles, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + articles, err := db.ListArticles(ctx, storage.ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, articles, 2) } @@ -204,7 +204,7 @@ func TestScanBlogRSSWithCategories(t *testing.T) { require.NoError(t, scanErr) require.Equal(t, 2, result.NewArticles) - articles, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + articles, err := db.ListArticles(ctx, storage.ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, articles, 2) diff --git a/internal/storage/database.go b/internal/storage/database.go index 9b3ab99..2531999 100644 --- a/internal/storage/database.go +++ b/internal/storage/database.go @@ -8,6 +8,7 @@ import ( "fmt" "os" "path/filepath" + "strings" "time" sq "github.com/Masterminds/squirrel" @@ -31,6 +32,11 @@ const ( // --since/--before lexicographic comparison in ListArticles is stable // regardless of the source precision. sqliteWriteLayout = "2006-01-02T15:04:05Z" + + // MaxListLimit is the maximum number of articles a caller may request via + // ArticleFilter.Limit. Callers that exceed it should be rejected by the + // CLI layer, not silently clamped. + MaxListLimit = 100 ) func DefaultDBPath() (string, error) { @@ -46,6 +52,17 @@ type Database struct { conn *sql.DB } +type ArticleFilter struct { + UnreadOnly bool + BlogID *int64 + Category *string + Since *time.Time + Before *time.Time + Search string + Limit int + Offset int +} + func OpenDatabase(ctx context.Context, path string) (*Database, error) { if path == "" { var err error @@ -59,11 +76,12 @@ func OpenDatabase(ctx context.Context, path string) (*Database, error) { return nil, err } - dsn := fmt.Sprintf("file:%s?_pragma=busy_timeout(5000)&_pragma=foreign_keys(1)", path) + dsn := fmt.Sprintf("file:%s?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=foreign_keys(1)", path) conn, err := sql.Open("sqlite", dsn) if err != nil { return nil, err } + conn.SetMaxOpenConns(1) db := &Database{path: path, conn: conn} if err := db.migrate(); err != nil { @@ -235,8 +253,8 @@ func (db *Database) AddArticle(ctx context.Context, article model.Article) (mode return article, err } result, err := sq.Insert("articles"). - Columns("blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories"). - Values(article.BlogID, article.Title, article.URL, formatTimePtr(article.PublishedDate), formatTimePtr(article.DiscoveredDate), article.IsRead, cats). + Columns("blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories", "description", "content"). + Values(article.BlogID, article.Title, article.URL, formatTimePtr(article.PublishedDate), formatTimePtr(article.DiscoveredDate), article.IsRead, cats, nullIfEmpty(article.Description), nullIfEmpty(article.Content)). RunWith(db.conn). ExecContext(ctx) if err != nil { @@ -260,7 +278,7 @@ func (db *Database) AddArticlesBulk(ctx context.Context, articles []model.Articl } insert := sq.Insert("articles"). - Columns("blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories") + Columns("blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories", "description", "content") for _, article := range articles { cats, err := categoriesToJSON(article.Categories) if err != nil { @@ -277,6 +295,8 @@ func (db *Database) AddArticlesBulk(ctx context.Context, articles []model.Articl formatTimePtr(article.DiscoveredDate), article.IsRead, cats, + nullIfEmpty(article.Description), + nullIfEmpty(article.Content), ) } @@ -295,7 +315,7 @@ func (db *Database) AddArticlesBulk(ctx context.Context, articles []model.Articl } func (db *Database) GetArticle(ctx context.Context, id int64) (*model.Article, error) { - row := sq.Select("id", "blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories"). + row := sq.Select("id", "blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories", "description", "content"). From("articles"). Where(sq.Eq{"id": id}). RunWith(db.conn). @@ -304,7 +324,7 @@ func (db *Database) GetArticle(ctx context.Context, id int64) (*model.Article, e } func (db *Database) GetArticleByURL(ctx context.Context, url string) (*model.Article, error) { - row := sq.Select("id", "blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories"). + row := sq.Select("id", "blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories", "description", "content"). From("articles"). Where(sq.Eq{"url": url}). RunWith(db.conn). @@ -371,27 +391,88 @@ func (db *Database) GetExistingArticleURLs(ctx context.Context, urls []string) ( return result, nil } -func (db *Database) ListArticles(ctx context.Context, unreadOnly bool, blogID *int64, category *string, since *time.Time, before *time.Time) ([]model.Article, error) { - query := sq.Select("id", "blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories"). +func (db *Database) ListArticles(ctx context.Context, filter ArticleFilter) ([]model.Article, error) { + if filter.Search != "" { + return db.searchArticles(ctx, filter) + } + + query := sq.Select("id", "blog_id", "title", "url", "published_date", "discovered_date", "is_read", "categories", "description", "content"). From("articles"). OrderBy("discovered_date DESC") - if unreadOnly { + if filter.UnreadOnly { query = query.Where(sq.Eq{"is_read": false}) } - if blogID != nil { - query = query.Where(sq.Eq{"blog_id": *blogID}) + if filter.BlogID != nil { + query = query.Where(sq.Eq{"blog_id": *filter.BlogID}) + } + if filter.Category != nil && *filter.Category != "" { + query = query.Where("EXISTS (SELECT 1 FROM json_each(categories) WHERE LOWER(json_each.value) = LOWER(?))", *filter.Category) + } + if filter.Since != nil { + query = query.Where(sq.GtOrEq{"published_date": filter.Since.UTC().Format(sqliteWriteLayout)}) + } + if filter.Before != nil { + query = query.Where(sq.Lt{"published_date": filter.Before.UTC().Format(sqliteWriteLayout)}) + } + if filter.Limit > 0 { + query = query.Limit(uint64(filter.Limit)) + } + if filter.Offset > 0 { + query = query.Offset(uint64(filter.Offset)) + } + + rows, err := query.RunWith(db.conn).QueryContext(ctx) + if err != nil { + return nil, err + } + defer func() { + if err := rows.Close(); err != nil { + fmt.Fprintf(os.Stderr, "close rows: %v\n", err) + } + }() + + var articles []model.Article + for rows.Next() { + article, err := scanArticle(rows) + if err != nil { + return nil, err + } + if article != nil { + articles = append(articles, *article) + } + } + return articles, rows.Err() +} + +func (db *Database) searchArticles(ctx context.Context, filter ArticleFilter) ([]model.Article, error) { + query := sq.Select("a.id", "a.blog_id", "a.title", "a.url", "a.published_date", "a.discovered_date", "a.is_read", "a.categories", "a.description", "a.content"). + From("articles a"). + Join("articles_fts f ON a.id = f.rowid"). + Where("articles_fts MATCH ?", escapeFTS5Query(filter.Search)). + OrderBy("f.rank") + + if filter.UnreadOnly { + query = query.Where(sq.Eq{"a.is_read": false}) + } + if filter.BlogID != nil { + query = query.Where(sq.Eq{"a.blog_id": *filter.BlogID}) + } + if filter.Category != nil && *filter.Category != "" { + query = query.Where("EXISTS (SELECT 1 FROM json_each(a.categories) WHERE LOWER(json_each.value) = LOWER(?))", *filter.Category) + } + if filter.Since != nil { + query = query.Where(sq.GtOrEq{"a.published_date": filter.Since.UTC().Format(sqliteWriteLayout)}) } - if category != nil && *category != "" { - // Categories are stored as a JSON string array. Use json_each() - // for exact element matching. - query = query.Where("EXISTS (SELECT 1 FROM json_each(categories) WHERE LOWER(json_each.value) = LOWER(?))", *category) + if filter.Before != nil { + query = query.Where(sq.Lt{"a.published_date": filter.Before.UTC().Format(sqliteWriteLayout)}) } - if since != nil { - query = query.Where(sq.GtOrEq{"published_date": since.UTC().Format(sqliteWriteLayout)}) + + if filter.Limit > 0 { + query = query.Limit(uint64(filter.Limit)) } - if before != nil { - query = query.Where(sq.Lt{"published_date": before.UTC().Format(sqliteWriteLayout)}) + if filter.Offset > 0 { + query = query.Offset(uint64(filter.Offset)) } rows, err := query.RunWith(db.conn).QueryContext(ctx) @@ -417,6 +498,20 @@ func (db *Database) ListArticles(ctx context.Context, unreadOnly bool, blogID *i return articles, rows.Err() } +// escapeFTS5Query prepares a user-supplied search string for use in an FTS5 +// MATCH clause. Words containing characters that are not valid in bare FTS5 +// tokens (single quotes, double quotes) are wrapped in double quotes so they +// are treated as phrase tokens rather than causing a syntax error. +func escapeFTS5Query(q string) string { + words := strings.Fields(q) + for i, w := range words { + if strings.ContainsAny(w, "'\"") { + words[i] = `"` + strings.ReplaceAll(w, `"`, `""`) + `"` + } + } + return strings.Join(words, " ") +} + func (db *Database) MarkArticleRead(ctx context.Context, id int64) (bool, error) { result, err := sq.Update("articles"). Set("is_read", true). @@ -433,6 +528,18 @@ func (db *Database) MarkArticleRead(ctx context.Context, id int64) (bool, error) return rows > 0, nil } +func (db *Database) MarkArticlesRead(ctx context.Context, ids []int64) error { + if len(ids) == 0 { + return nil + } + _, err := sq.Update("articles"). + Set("is_read", true). + Where(sq.Eq{"id": ids}). + RunWith(db.conn). + ExecContext(ctx) + return err +} + func (db *Database) MarkArticleUnread(ctx context.Context, id int64) (bool, error) { result, err := sq.Update("articles"). Set("is_read", false). @@ -492,8 +599,10 @@ func scanArticle(scanner interface{ Scan(dest ...any) error }) (*model.Article, discovered sql.NullString isRead bool categories sql.NullString + description sql.NullString + content sql.NullString ) - if err := scanner.Scan(&id, &blogID, &title, &url, &publishedDate, &discovered, &isRead, &categories); err != nil { + if err := scanner.Scan(&id, &blogID, &title, &url, &publishedDate, &discovered, &isRead, &categories, &description, &content); err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, nil } @@ -506,12 +615,14 @@ func scanArticle(scanner interface{ Scan(dest ...any) error }) (*model.Article, } article := &model.Article{ - ID: id, - BlogID: blogID, - Title: title, - URL: url, - IsRead: isRead, - Categories: cats, + ID: id, + BlogID: blogID, + Title: title, + URL: url, + IsRead: isRead, + Categories: cats, + Description: description.String, + Content: content.String, } if publishedDate.Valid { if parsed, err := parseTime(publishedDate.String); err == nil { diff --git a/internal/storage/database_test.go b/internal/storage/database_test.go index d9877ad..27345b6 100644 --- a/internal/storage/database_test.go +++ b/internal/storage/database_test.go @@ -42,7 +42,7 @@ func TestDatabaseCreatesFileAndCRUD(t *testing.T) { require.NoError(t, err, "add articles bulk") require.Equal(t, 2, count) - list, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + list, err := db.ListArticles(ctx, ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, list, 2) @@ -197,17 +197,17 @@ func TestListArticlesFiltersAndOrdering(t *testing.T) { _, err = db.MarkArticleRead(ctx, first.ID) require.NoError(t, err, "mark read") - all, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + all, err := db.ListArticles(ctx, ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, all, 3) require.Equal(t, second.ID, all[0].ID, "expected newest article first") - unread, err := db.ListArticles(ctx, true, nil, nil, nil, nil) + unread, err := db.ListArticles(ctx, ArticleFilter{UnreadOnly: true}) require.NoError(t, err, "list unread") require.Len(t, unread, 2) blogID := blogB.ID - filtered, err := db.ListArticles(ctx, false, &blogID, nil, nil, nil) + filtered, err := db.ListArticles(ctx, ArticleFilter{BlogID: &blogID}) require.NoError(t, err, "list by blog") require.Len(t, filtered, 1) require.Equal(t, blogB.ID, filtered[0].BlogID) @@ -238,7 +238,7 @@ func TestBulkInsertDuplicateRollbackAndEmpty(t *testing.T) { _, err = db.AddArticlesBulk(ctx, dupArticles) require.Error(t, err, "expected bulk insert to fail on duplicate url") - articles, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + articles, err := db.ListArticles(ctx, ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, articles, 1, "expected rollback on duplicate") } @@ -353,38 +353,38 @@ func TestListArticlesFilterByCategory(t *testing.T) { // Filter by "Go" - should return only the Go article cat := "Go" - goArticles, err := db.ListArticles(ctx, false, nil, &cat, nil, nil) + goArticles, err := db.ListArticles(ctx, ArticleFilter{Category: &cat}) require.NoError(t, err, "list by category Go") require.Len(t, goArticles, 1) require.Equal(t, "Go Article", goArticles[0].Title) // Filter by "Programming" - should return both categorized articles cat = "Programming" - progArticles, err := db.ListArticles(ctx, false, nil, &cat, nil, nil) + progArticles, err := db.ListArticles(ctx, ArticleFilter{Category: &cat}) require.NoError(t, err, "list by category Programming") require.Len(t, progArticles, 2) // No filter - should return all 3 - all, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + all, err := db.ListArticles(ctx, ArticleFilter{}) require.NoError(t, err, "list all") require.Len(t, all, 3) // Case-insensitive match - "go" should match "Go" cat = "go" - goLower, err := db.ListArticles(ctx, false, nil, &cat, nil, nil) + goLower, err := db.ListArticles(ctx, ArticleFilter{Category: &cat}) require.NoError(t, err, "list by category go (lowercase)") require.Len(t, goLower, 1) require.Equal(t, "Go Article", goLower[0].Title) // Case-insensitive match - "PROGRAMMING" should match "Programming" cat = "PROGRAMMING" - progUpper, err := db.ListArticles(ctx, false, nil, &cat, nil, nil) + progUpper, err := db.ListArticles(ctx, ArticleFilter{Category: &cat}) require.NoError(t, err, "list by category PROGRAMMING (uppercase)") require.Len(t, progUpper, 2) // Empty string category should return all empty := "" - allEmpty, err := db.ListArticles(ctx, false, nil, &empty, nil, nil) + allEmpty, err := db.ListArticles(ctx, ArticleFilter{Category: &empty}) require.NoError(t, err, "list with empty category") require.Len(t, allEmpty, 3) } @@ -405,7 +405,7 @@ func TestBulkInsertWithCategories(t *testing.T) { require.NoError(t, err, "bulk insert") require.Equal(t, 2, count) - list, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + list, err := db.ListArticles(ctx, ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, list, 2) @@ -450,14 +450,14 @@ func TestListArticlesFilterByDate(t *testing.T) { require.NoError(t, err, "add article without date") t.Run("without filters returns all articles", func(t *testing.T) { - articles, err := db.ListArticles(ctx, false, nil, nil, nil, nil) + articles, err := db.ListArticles(ctx, ArticleFilter{}) require.NoError(t, err, "list articles") require.Len(t, articles, 4, "should return all articles including no-date article") }) t.Run("since filter inclusive", func(t *testing.T) { since := time.Date(2024, 1, 15, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, &since, nil) + articles, err := db.ListArticles(ctx, ArticleFilter{Since: &since}) require.NoError(t, err, "list articles with since filter") require.Len(t, articles, 2, "should return articles on or after since date (Article2 and Article3)") titles := []string{articles[0].Title, articles[1].Title} @@ -467,7 +467,7 @@ func TestListArticlesFilterByDate(t *testing.T) { t.Run("before filter exclusive", func(t *testing.T) { before := time.Date(2024, 1, 15, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, nil, &before) + articles, err := db.ListArticles(ctx, ArticleFilter{Before: &before}) require.NoError(t, err, "list articles with before filter") require.Len(t, articles, 1, "should return articles before date (only Article1)") require.Equal(t, "Article1", articles[0].Title, "should only include Article1 before before-date") @@ -476,7 +476,7 @@ func TestListArticlesFilterByDate(t *testing.T) { t.Run("combined filters", func(t *testing.T) { since := time.Date(2024, 1, 10, 0, 0, 0, 0, time.UTC) before := time.Date(2024, 2, 1, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, &since, &before) + articles, err := db.ListArticles(ctx, ArticleFilter{Since: &since, Before: &before}) require.NoError(t, err, "list articles with combined filters") require.Len(t, articles, 1, "should return only Article2 in range") require.Equal(t, "Article2", articles[0].Title, "should only include Article2") @@ -484,7 +484,7 @@ func TestListArticlesFilterByDate(t *testing.T) { t.Run("nil published date excluded from filters", func(t *testing.T) { since := time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, &since, nil) + articles, err := db.ListArticles(ctx, ArticleFilter{Since: &since}) require.NoError(t, err, "list articles with since filter") require.Len(t, articles, 3, "should exclude no-date article") @@ -495,14 +495,14 @@ func TestListArticlesFilterByDate(t *testing.T) { t.Run("after all dates", func(t *testing.T) { since := time.Date(2024, 3, 1, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, &since, nil) + articles, err := db.ListArticles(ctx, ArticleFilter{Since: &since}) require.NoError(t, err, "list articles with since filter after all dates") require.Empty(t, articles, "should return empty result") }) t.Run("before all dates", func(t *testing.T) { before := time.Date(2023, 1, 1, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, nil, &before) + articles, err := db.ListArticles(ctx, ArticleFilter{Before: &before}) require.NoError(t, err, "list articles with before filter before all dates") require.Empty(t, articles, "should return empty result") }) @@ -580,7 +580,7 @@ func TestDateFilterRespectsTimezoneEquivalence(t *testing.T) { require.NoError(t, err, "add article") since := time.Date(2024, 1, 15, 0, 0, 0, 0, time.UTC) - articles, err := db.ListArticles(ctx, false, nil, nil, &since, nil) + articles, err := db.ListArticles(ctx, ArticleFilter{Since: &since}) require.NoError(t, err, "list articles") require.Empty(t, articles, "JST article published before UTC midnight Jan 15 should be excluded") } @@ -670,6 +670,134 @@ func TestMigrationPreservesUnparseableTimestamps(t *testing.T) { require.Equal(t, garbage, stored, "unparseable timestamp should be preserved verbatim") } +func TestListArticlesFilterBySearch(t *testing.T) { + ctx := context.Background() + db := openTestDB(t) + defer func() { require.NoError(t, db.Close()) }() + + blog, err := db.AddBlog(ctx, model.Blog{Name: "Test", URL: "https://example.com"}) + require.NoError(t, err, "add blog") + + _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "Golang concurrency patterns", URL: "https://example.com/go-concurrency"}) + require.NoError(t, err, "add article") + _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "Rust memory safety", URL: "https://example.com/rust-memory"}) + require.NoError(t, err, "add article") + _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "Python async pitfalls", URL: "https://example.com/python-async"}) + require.NoError(t, err, "add article") + + t.Run("search by single word", func(t *testing.T) { + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "golang"}) + require.NoError(t, err) + require.Len(t, articles, 1) + require.Equal(t, "Golang concurrency patterns", articles[0].Title) + }) + + t.Run("search is case-insensitive", func(t *testing.T) { + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "GOLANG"}) + require.NoError(t, err) + require.Len(t, articles, 1) + require.Equal(t, "Golang concurrency patterns", articles[0].Title) + }) + + t.Run("search with no matches returns empty", func(t *testing.T) { + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "nonexistent"}) + require.NoError(t, err) + require.Empty(t, articles) + }) + + t.Run("search combined with blog filter", func(t *testing.T) { + blogID := blog.ID + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "rust", BlogID: &blogID}) + require.NoError(t, err) + require.Len(t, articles, 1) + require.Equal(t, "Rust memory safety", articles[0].Title) + }) + + t.Run("search with apostrophe does not need escaping", func(t *testing.T) { + _, err := db.AddArticle(ctx, model.Article{ + BlogID: blog.ID, + Title: "Don't fear the goroutine", + URL: "https://example.com/dont", + Description: "You don't need to worry about goroutines", + }) + require.NoError(t, err, "add article with apostrophe") + + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "don't"}) + require.NoError(t, err) + require.Len(t, articles, 1, "should find article with apostrophe in query") + require.Equal(t, "Don't fear the goroutine", articles[0].Title) + }) + + t.Run("search matches body text", func(t *testing.T) { + _, err := db.AddArticle(ctx, model.Article{ + BlogID: blog.ID, + Title: "Deep dive into channels", + URL: "https://example.com/channels", + Description: "This article discusses golang concurrency channels in detail", + }) + require.NoError(t, err, "add article with body text") + + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "channels"}) + require.NoError(t, err) + require.Len(t, articles, 1, "should find article by body text match") + require.Equal(t, "Deep dive into channels", articles[0].Title) + }) + + t.Run("search combined with unread filter", func(t *testing.T) { + golangArticles, err := db.ListArticles(ctx, ArticleFilter{Search: "golang"}) + require.NoError(t, err) + require.NotEmpty(t, golangArticles, "expected at least one article matching 'golang'") + + for _, a := range golangArticles { + _, err = db.MarkArticleRead(ctx, a.ID) + require.NoError(t, err) + } + + articles, err := db.ListArticles(ctx, ArticleFilter{Search: "golang", UnreadOnly: true}) + require.NoError(t, err) + require.Empty(t, articles, "expected no unread articles matching 'golang' after marking all matching articles as read") + }) +} + +func TestListArticlesWithLimit(t *testing.T) { + ctx := context.Background() + db := openTestDB(t) + defer func() { require.NoError(t, db.Close()) }() + + blog, err := db.AddBlog(ctx, model.Blog{Name: "Test", URL: "https://example.com"}) + require.NoError(t, err, "add blog") + + t1 := time.Date(2024, 1, 1, 10, 0, 0, 0, time.UTC) + t2 := time.Date(2024, 1, 2, 10, 0, 0, 0, time.UTC) + t3 := time.Date(2024, 1, 3, 10, 0, 0, 0, time.UTC) + + _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "First", URL: "https://example.com/1", DiscoveredDate: &t1}) + require.NoError(t, err, "add article 1") + _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "Second", URL: "https://example.com/2", DiscoveredDate: &t2}) + require.NoError(t, err, "add article 2") + _, err = db.AddArticle(ctx, model.Article{BlogID: blog.ID, Title: "Third", URL: "https://example.com/3", DiscoveredDate: &t3}) + require.NoError(t, err, "add article 3") + + t.Run("limit returns fewer results", func(t *testing.T) { + articles, err := db.ListArticles(ctx, ArticleFilter{Limit: 2}) + require.NoError(t, err) + require.Len(t, articles, 2) + }) + + t.Run("limit of 0 returns all", func(t *testing.T) { + articles, err := db.ListArticles(ctx, ArticleFilter{Limit: 0}) + require.NoError(t, err) + require.Len(t, articles, 3) + }) + + t.Run("limit with offset", func(t *testing.T) { + articles, err := db.ListArticles(ctx, ArticleFilter{Limit: 1, Offset: 1}) + require.NoError(t, err) + require.Len(t, articles, 1) + require.Equal(t, "Second", articles[0].Title) + }) +} + func openTestDB(t *testing.T) *Database { t.Helper() ctx := context.Background() diff --git a/internal/storage/migrations/000004_add_description_content.down.sql b/internal/storage/migrations/000004_add_description_content.down.sql new file mode 100644 index 0000000..d16028f --- /dev/null +++ b/internal/storage/migrations/000004_add_description_content.down.sql @@ -0,0 +1,2 @@ +ALTER TABLE articles DROP COLUMN content; +ALTER TABLE articles DROP COLUMN description; diff --git a/internal/storage/migrations/000004_add_description_content.up.sql b/internal/storage/migrations/000004_add_description_content.up.sql new file mode 100644 index 0000000..aa343ff --- /dev/null +++ b/internal/storage/migrations/000004_add_description_content.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE articles ADD COLUMN description TEXT; +ALTER TABLE articles ADD COLUMN content TEXT; diff --git a/internal/storage/migrations/000005_add_search.down.sql b/internal/storage/migrations/000005_add_search.down.sql new file mode 100644 index 0000000..0795d50 --- /dev/null +++ b/internal/storage/migrations/000005_add_search.down.sql @@ -0,0 +1,4 @@ +DROP TRIGGER IF EXISTS articles_ai; +DROP TRIGGER IF EXISTS articles_ad; +DROP TRIGGER IF EXISTS articles_au; +DROP TABLE IF EXISTS articles_fts; diff --git a/internal/storage/migrations/000005_add_search.up.sql b/internal/storage/migrations/000005_add_search.up.sql new file mode 100644 index 0000000..51c56bb --- /dev/null +++ b/internal/storage/migrations/000005_add_search.up.sql @@ -0,0 +1,25 @@ +CREATE VIRTUAL TABLE articles_fts USING fts5( + title, description, content, + content='articles', content_rowid='rowid' +); + +INSERT INTO articles_fts(rowid, title, description, content) + SELECT id, title, COALESCE(description, ''), COALESCE(content, '') + FROM articles; + +CREATE TRIGGER articles_ai AFTER INSERT ON articles BEGIN + INSERT INTO articles_fts(rowid, title, description, content) + VALUES (new.id, new.title, COALESCE(new.description, ''), COALESCE(new.content, '')); +END; + +CREATE TRIGGER articles_ad AFTER DELETE ON articles BEGIN + INSERT INTO articles_fts(articles_fts, rowid, title, description, content) + VALUES ('delete', old.id, old.title, COALESCE(old.description, ''), COALESCE(old.content, '')); +END; + +CREATE TRIGGER articles_au AFTER UPDATE ON articles BEGIN + INSERT INTO articles_fts(articles_fts, rowid, title, description, content) + VALUES ('delete', old.id, old.title, COALESCE(old.description, ''), COALESCE(old.content, '')); + INSERT INTO articles_fts(rowid, title, description, content) + VALUES (new.id, new.title, COALESCE(new.description, ''), COALESCE(new.content, '')); +END;