package cli import ( "crypto/rand" "errors" "fmt" "log/slog" "math/big" "net" "net/mail" "os" "strings" "text/tabwriter" "postern/internal/db" "github.com/spf13/cobra" "golang.org/x/crypto/bcrypt" ) func (c *Command) addUser(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true // suppress runtime error help text ctx := cmd.Context() username := args[0] if exists, err := c.db.UserExists(ctx, username); err != nil || exists { if err != nil { return err } return fmt.Errorf("user %q already exists", username) } password, err := generateSafeToken(16) if err != nil { return err } hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return err } tx, err := c.db.GetWriteTx(ctx) if err != nil { return fmt.Errorf("beginning transaction: %w", err) } userID, err := db.InsertUser(ctx, tx, username, hash[:]) if err != nil { txErr := tx.Rollback() if txErr != nil { slog.Error("rolling back transaction", "error", err) } return fmt.Errorf("inserting user %q: %w", username, err) } if err := db.InsertDefaultMailboxes(ctx, tx, userID); err != nil { txErr := tx.Rollback() if txErr != nil { slog.Error("rolling back transaction", "error", err) } return fmt.Errorf("inserting default mailboxes: %w", err) } if err := tx.Commit(); err != nil { return fmt.Errorf("committing transaction: %w", err) } // subscribe user to all mailboxes allMailboxes, err := c.db.GetUserMailboxes(ctx, userID) if err != nil { return err } for _, mb := range allMailboxes { err = c.db.Subscribe(ctx, userID, mb.ID, mb.Name) if err != nil { return fmt.Errorf("subscribing mailbox %q: %w", mb.Name, err) } slog.Debug("subscribed mailbox", "mailbox", mb.Name) } _, _ = fmt.Fprintf(cmd.OutOrStdout(), "User %q created. Your password will only be shown once.\n", username) _, _ = fmt.Fprintf(cmd.OutOrStdout(), "Password: %s\n", password) return nil } func (c *Command) removeUser(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true // suppress runtime error help text ctx := cmd.Context() username := args[0] if exists, err := c.db.UserExists(ctx, username); err != nil || !exists { if err != nil { return err } return fmt.Errorf("user %q does not exist", username) } err := c.db.RemoveUser(ctx, username) if err != nil { return fmt.Errorf("removing user %q: %w", username, err) } _, _ = fmt.Fprintf(cmd.OutOrStdout(), "User %q removed.\n", username) return nil } func (c *Command) listUsers(cmd *cobra.Command, _ []string) error { cmd.SilenceUsage = true users, err := c.db.ListUsers(cmd.Context()) if err != nil { fmt.Println("Error listing users") return err } w := tabwriter.NewWriter(os.Stdout, 1, 1, 1, ' ', 0) _, _ = fmt.Fprintln(w, "Name") for _, user := range users { _, _ = fmt.Fprintf(w, "%s\n", user.Name) } _ = w.Flush() return nil } func (c *Command) addUserAddress(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true // suppress runtime error help text ctx := cmd.Context() username := args[0] address := args[1] parsedAddress, err := mail.ParseAddress(address) if err != nil { fmt.Printf("Could not parse address: %s\n", address) return err } i := strings.LastIndex(parsedAddress.Address, "@") if i == -1 { fmt.Printf("invalid address does not have a domain part: %s\n", parsedAddress.Address) return errors.New("invalid address") } err = c.db.InsertUserAddress(ctx, username, parsedAddress.Address) if err != nil { return err } fmt.Printf("address <%s> added for %q\n", address, username) return nil } func (c *Command) addDKIM(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true selector := args[0] domain := args[1] record, err := c.dkim.Add(cmd.Context(), selector, domain) if err != nil { return err } fmt.Println("Create a DNS TXT record for", domainkeyName(selector, domain)) fmt.Println("Remember to enable this selector once it was made available through DNS. (dkim enable)") fmt.Println(record) return nil } func (c *Command) listDKIM(cmd *cobra.Command, _ []string) error { cmd.SilenceUsage = true // suppress runtime error help text dkims, err := c.dkim.List(cmd.Context()) if err != nil { fmt.Println("Error list DKIM") return err } if len(dkims) == 0 { fmt.Println("No DKIM records found") return nil } w := tabwriter.NewWriter(os.Stdout, 1, 1, 1, ' ', 0) _, _ = fmt.Fprintln(w, "Identifier\tKey Type\tEnabled") for _, dkim := range dkims { _, _ = fmt.Fprintf(w, "%s\t%s\t%v\n", domainkeyName(dkim.Selector, dkim.Domain), dkim.KeyType, dkim.Enabled) } _ = w.Flush() return nil } func (c *Command) getDKIM(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true selector := args[0] domain := args[1] record, err := c.dkim.PublicDNSRecord(cmd.Context(), selector, domain) if err != nil { return err } fmt.Println("Create a DNS TXT record for", domainkeyName(selector, domain)) fmt.Println("Remember to enable this selector once it was made available through DNS. (dkim enable)") fmt.Println(record) return nil } func (c *Command) setDKIMEnabled(cmd *cobra.Command, args []string, enable bool) error { cmd.SilenceUsage = true selector := args[0] domain := args[1] if enable { force, err := cmd.Flags().GetBool("force") if err != nil { return err } if !force { s := domainkeyName(selector, domain) records, err := net.LookupTXT(s) if err != nil { fmt.Println("resolving TXT records: %w", err) } expected, err := c.dkim.PublicDNSRecord(cmd.Context(), selector, domain) if err != nil { return err } if expected != strings.Join(records, "") { fmt.Println("The DKIM record in DNS does not match the expected record.") fmt.Println("To enable this record anyway, use '--force'.") fmt.Println("Expected:") fmt.Println(expected) fmt.Println("Actual:") fmt.Println(strings.Join(records, "")) return errors.New("DKIM record does not match the expected record") } } err = c.dkim.SetEnabled(cmd.Context(), selector, domain, true) if err != nil { return err } fmt.Printf("Enabling DKIM selector %q for domain %q\n", selector, domain) } else { if err := c.dkim.SetEnabled(cmd.Context(), selector, domain, false); err != nil { return err } fmt.Printf("Disabling DKIM selector %q for domain %q\n", selector, domain) } return nil } func (c *Command) removeDKIM(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true selector := args[0] domain := args[1] return c.dkim.Remove(cmd.Context(), selector, domain) } func (c *Command) listSMTPAuth(cmd *cobra.Command, _ []string) error { cmd.SilenceUsage = true smtpAuths, err := c.db.ListSMTPAuthUsers(cmd.Context()) if err != nil { fmt.Println("Error listing SMTP auth users") return err } w := tabwriter.NewWriter(os.Stdout, 1, 1, 1, ' ', 0) _, _ = fmt.Fprintln(w, "Name") for _, auth := range smtpAuths { _, _ = fmt.Fprintf(w, "%s\n", auth.Name) } _ = w.Flush() return nil } func (c *Command) addSMTPAuth(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true // suppress runtime error help text ctx := cmd.Context() name := args[0] if exists, err := c.db.SMTPAuthNameExists(ctx, name); err != nil || exists { if err != nil { return err } return fmt.Errorf("user %q already exists", name) } password, err := generateSafeToken(16) if err != nil { return err } hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return err } err = c.db.InsertSMTPAuthUser(ctx, name, hash[:]) if err != nil { return fmt.Errorf("inserting smtp auth %q: %w", name, err) } _, _ = fmt.Fprintf(cmd.OutOrStdout(), "SMTP Auth %q created. The password will only be shown once.\n", name) _, _ = fmt.Fprintf(cmd.OutOrStdout(), "Password: %s\n", password) return nil } func (c *Command) removeSMTPAuth(cmd *cobra.Command, args []string) error { cmd.SilenceUsage = true // suppress runtime error help text ctx := cmd.Context() name := args[0] if exists, err := c.db.SMTPAuthNameExists(ctx, name); err != nil || !exists { if err != nil { return err } return fmt.Errorf("user %q does not exist", name) } err := c.db.RemoveSMTPAuthUser(ctx, name) if err != nil { return err } _, _ = fmt.Fprintf(cmd.OutOrStdout(), "User %q removed.\n", name) return nil } func domainkeyName(selector, domain string) string { if !strings.HasSuffix(domain, ".") { domain = domain + "." } return fmt.Sprintf("%s._domainkey.%s", selector, domain) } // Define the character set: // - Numbers: 2-9 (No 0, 1) // - Uppercase: A-Z (No I, O) // - Lowercase: a-z (No l, o) const letterBytes = "23456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz" func generateSafeToken(length int) (string, error) { if length <= 0 { return "", fmt.Errorf("length must be greater than 0") } ret := make([]byte, length) // We calculate the max index based on the charset length charsetLen := big.NewInt(int64(len(letterBytes))) for i := range length { // rand.Int returns a uniform random value in [0, max). // It automatically handles modulo bias, ensuring an even distribution. num, err := rand.Int(rand.Reader, charsetLen) if err != nil { return "", err } // Use the random number as an index to pick a character ret[i] = letterBytes[num.Int64()] } return string(ret), nil }