package command import ( "cmp" "context" "fmt" "os" "os/signal" "slices" "strconv" "strings" "syscall" "time" "github.com/StevanFreeborn/gator/internal/database" "github.com/StevanFreeborn/gator/internal/rss" "github.com/StevanFreeborn/gator/internal/state" "github.com/google/uuid" ) type Command struct { Name string Handler CommandHandler } type CommandHandler func(s *state.State) error func newCommand(name string, handler CommandHandler) *Command { return &Command{ Name: name, Handler: handler, } } type CommandRegistry struct { commands map[string]*Command } func (c *CommandRegistry) register(cmd *Command) error { c.commands[cmd.Name] = cmd return nil } func NewRegistry() *CommandRegistry { cr := CommandRegistry{ commands: map[string]*Command{}, } commands := []*Command{ loginCommand(), registerCommand(), resetCommand(), usersCommand(), aggCommand(), addFeedCommand(), feedsCommand(), followCommand(), followingCommand(), unfollowCommand(), browseCommand(), } for _, cmd := range commands { cr.register(cmd) } return &cr } func (c *CommandRegistry) RunCommand(cmdName string, s *state.State) error { cmd, found := c.commands[cmdName] if !found { return fmt.Errorf("No '%s' registered", cmdName) } return cmd.Handler(s) } func loginCommand() *Command { return newCommand("login", func(s *state.State) error { if len(s.Arguments) == 0 { return fmt.Errorf("Did not receive expected username argument") } username := s.Arguments[0] _, err := s.Database.GetUserByName(context.Background(), username) if err != nil { return fmt.Errorf("Failed to login") } err = s.Config.SetUser(username) if err != nil { return err } fmt.Printf("Current user set to '%s'\n", username) return nil }) } func registerCommand() *Command { return newCommand("register", func(s *state.State) error { if len(s.Arguments) == 0 { return fmt.Errorf("Did not receive expected username argument") } createUserParams := database.CreateUserParams{ ID: uuid.New(), CreatedAt: time.Now(), UpdatedAt: time.Now(), Name: s.Arguments[0], } createdUser, err := s.Database.CreateUser(context.Background(), createUserParams) if err != nil { return err } s.Config.SetUser(createdUser.Name) fmt.Printf("Successfully registered user '%s':\n", createdUser.Name) fmt.Printf(" Id => %s\n", createdUser.ID) fmt.Printf(" CreatedAt => %s\n", createdUser.CreatedAt) fmt.Printf(" UpdatedAt => %s\n", createdUser.UpdatedAt) fmt.Printf("Current user set to user '%s':\n", createdUser.Name) return nil }) } func resetCommand() *Command { return newCommand("reset", func(s *state.State) error { err := s.Database.DeleteAllUsers(context.Background()) if err != nil { return err } fmt.Println("Successfully delete all users") return nil }) } func usersCommand() *Command { return newCommand("users", func(s *state.State) error { users, err := s.Database.GetAllUsers(context.Background()) if err != nil { return err } slices.SortFunc(users, func(a, b database.User) int { return cmp.Compare(a.Name, b.Name) }) for _, user := range users { msg := "* %s" if user.Name == s.Config.CurrentUserName { msg += " (current)" } fmt.Printf(msg+"\n", user.Name) } return nil }) } func aggCommand() *Command { return newCommand("agg", func(s *state.State) error { if len(s.Arguments) == 0 { return fmt.Errorf("Did not receive expected time between requests argument") } timeBetweenRequests := s.Arguments[0] validDuration, err := time.ParseDuration(timeBetweenRequests) if err != nil { return fmt.Errorf("Time between requests argument '%s' not valid duration string", timeBetweenRequests) } ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer cancel() restore, err := enableKeypressExit(cancel) if err != nil { return err } defer restore() ticker := time.NewTicker(validDuration) defer ticker.Stop() rss.ScrapeNextFeed(ctx, s) fetchFeeds: for { select { case <-ctx.Done(): fmt.Print("Stopping feed agg command\r\n") break fetchFeeds case <-ticker.C: rss.ScrapeNextFeed(ctx, s) } } return nil }) } func addFeedCommand() *Command { return newCommand("addfeed", requiresLoggedInUser(func(s *state.State, currentUser database.User) error { if len(s.Arguments) < 2 { return fmt.Errorf("Did not receive expected feed name and url") } feedName := s.Arguments[0] feedUrl := s.Arguments[1] if strings.TrimSpace(feedName) == "" { return fmt.Errorf("Feed name cannot be empty") } canonicalFeedUrl, err := normalizeFeedURL(feedUrl) if err != nil { return fmt.Errorf("Feed url '%s' is not a valid url", feedUrl) } _, err = s.Database.GetFeedByUrl(context.Background(), canonicalFeedUrl) if err == nil { return fmt.Errorf("Feed with url '%s' already exists; use 'follow ' to follow it", canonicalFeedUrl) } createFeedParams := database.CreateFeedParams{ ID: uuid.New(), UserID: currentUser.ID, Name: feedName, Url: canonicalFeedUrl, CreatedAt: time.Now(), UpdatedAt: time.Now(), } createdFeed, err := s.Database.CreateFeed(context.Background(), createFeedParams) if err != nil { return err } createFollowParams := database.CreateFollowParams{ ID: uuid.New(), CreatedAt: time.Now(), UpdatedAt: time.Now(), UserID: currentUser.ID, FeedID: createdFeed.ID, } _, err = s.Database.CreateFollow(context.Background(), createFollowParams) if err != nil { return err } fmt.Printf("Successfully added and followed feed '%s' with url '%s'\n", createdFeed.Name, createdFeed.Url) fmt.Printf(" Id => %s\n", createdFeed.ID) fmt.Printf(" UserId => %s\n", createdFeed.UserID) fmt.Printf(" CreatedAt => %s\n", createdFeed.CreatedAt) fmt.Printf(" UpdatedAt => %s\n", createdFeed.UpdatedAt) return nil })) } func feedsCommand() *Command { return newCommand("feeds", func(s *state.State) error { feeds, err := s.Database.GetAllFeeds(context.Background()) if err != nil { return err } for _, feed := range feeds { feedsUserName := "Unknown" user, err := s.Database.GetUserById(context.Background(), feed.UserID) if err == nil { feedsUserName = user.Name } fmt.Printf("* %s [%s] (%s)\n", feed.Name, feed.Url, feedsUserName) } return nil }) } func followCommand() *Command { return newCommand("follow", requiresLoggedInUser(func(s *state.State, currentUser database.User) error { if len(s.Arguments) == 0 { return fmt.Errorf("Did not receive expected url argument") } feedIdentifier := s.Arguments[0] feed, err := resolveFeedByNameOrURL(s, feedIdentifier) if err != nil { return err } createFollowParams := database.CreateFollowParams{ ID: uuid.New(), CreatedAt: time.Now(), UpdatedAt: time.Now(), FeedID: feed.ID, UserID: currentUser.ID, } createdFollow, err := s.Database.CreateFollow(context.Background(), createFollowParams) if err != nil { return fmt.Errorf("Unable to follow feed '%s'", feed.Name) } fmt.Printf("'%s' successfully followed feed '%s'", createdFollow.UserName, createdFollow.FeedName) return nil })) } func followingCommand() *Command { return newCommand("following", requiresLoggedInUser(func(s *state.State, currentUser database.User) error { follows, err := s.Database.GetFollowsForUser(context.Background(), currentUser.ID) if err != nil { return fmt.Errorf("Unable to find follows for user") } for _, follow := range follows { fmt.Printf("* %s\n", follow.FeedName) } return nil })) } func unfollowCommand() *Command { return newCommand("unfollow", requiresLoggedInUser(func(s *state.State, currentUser database.User) error { if len(s.Arguments) == 0 { return fmt.Errorf("Did not receive expected feed url argument") } feedIdentifier := s.Arguments[0] feed, err := resolveFeedByNameOrURL(s, feedIdentifier) if err != nil { return err } deleteFollowParams := database.DeleteFollowForUserParams{ UserID: currentUser.ID, FeedID: feed.ID, } err = s.Database.DeleteFollowForUser(context.Background(), deleteFollowParams) if err != nil { return fmt.Errorf("Unable to unfollow feed '%s'", feed.Name) } fmt.Printf("Successfully unfollowed feed '%s'", feed.Name) return nil })) } func browseCommand() *Command { return newCommand("browse", requiresLoggedInUser(func(s *state.State, currentUser database.User) error { limit := int32(2) if len(s.Arguments) > 0 { parsedLimit, err := strconv.ParseInt(s.Arguments[0], 10, 32) if err == nil { limit = int32(parsedLimit) } } getUserPostsParams := database.GetPostsForUserParams{ UserID: currentUser.ID, Limit: limit, } posts, err := s.Database.GetPostsForUser(context.Background(), getUserPostsParams) if err != nil { return err } for _, post := range posts { fmt.Printf("* %s - %s\n", post.FeedName, post.Title) } return nil })) }