Build the Go API and React UI for managing domain denylist/allowlist/recipient deny and spam settings, plus shared lists that import/apply across domains with wildcard compaction, entry verification, and Migadu-friendly list encoding (including per-domain recipient filtering and rejection bisect on apply).
470 lines
12 KiB
Go
470 lines
12 KiB
Go
package store
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/squid/MigaduAdmin/internal/lists"
|
|
_ "modernc.org/sqlite"
|
|
)
|
|
|
|
var (
|
|
ErrNotFound = errors.New("not found")
|
|
ErrEmailTaken = errors.New("email already registered")
|
|
ErrInvalidInput = errors.New("invalid input")
|
|
ErrDomainExists = errors.New("domain already managed")
|
|
)
|
|
|
|
type Role string
|
|
|
|
const (
|
|
RoleAdmin Role = "admin"
|
|
RoleUser Role = "user"
|
|
)
|
|
|
|
type User struct {
|
|
ID int64 `json:"id"`
|
|
Email string `json:"email"`
|
|
DisplayName string `json:"display_name"`
|
|
Role Role `json:"role"`
|
|
PasswordHash string `json:"-"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
}
|
|
|
|
type Session struct {
|
|
ID string
|
|
UserID int64
|
|
ExpiresAt time.Time
|
|
CreatedAt time.Time
|
|
}
|
|
|
|
type ManagedDomain struct {
|
|
Name string `json:"name"`
|
|
AddedAt time.Time `json:"added_at"`
|
|
AddedBy *int64 `json:"added_by,omitempty"`
|
|
}
|
|
|
|
// SharedLists holds denylist/allowlist entries applied across managed domains.
|
|
type SharedLists struct {
|
|
SenderDenylist []string `json:"sender_denylist"`
|
|
SenderAllowlist []string `json:"sender_allowlist"`
|
|
RecipientDenylist []string `json:"recipient_denylist"`
|
|
}
|
|
|
|
type Store struct {
|
|
db *sql.DB
|
|
}
|
|
|
|
func Open(path string) (*Store, error) {
|
|
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
|
return nil, fmt.Errorf("create database directory: %w", err)
|
|
}
|
|
|
|
db, err := sql.Open("sqlite", path)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("open sqlite: %w", err)
|
|
}
|
|
db.SetMaxOpenConns(1)
|
|
|
|
store := &Store{db: db}
|
|
if err := store.migrate(); err != nil {
|
|
_ = db.Close()
|
|
return nil, err
|
|
}
|
|
return store, nil
|
|
}
|
|
|
|
func (s *Store) Close() error {
|
|
return s.db.Close()
|
|
}
|
|
|
|
func (s *Store) migrate() error {
|
|
const schema = `
|
|
CREATE TABLE IF NOT EXISTS users (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
email TEXT NOT NULL UNIQUE COLLATE NOCASE,
|
|
display_name TEXT NOT NULL DEFAULT '',
|
|
role TEXT NOT NULL,
|
|
password_hash TEXT NOT NULL DEFAULT '',
|
|
created_at TEXT NOT NULL
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS sessions (
|
|
id TEXT PRIMARY KEY,
|
|
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
expires_at TEXT NOT NULL,
|
|
created_at TEXT NOT NULL
|
|
);
|
|
|
|
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions(user_id);
|
|
CREATE INDEX IF NOT EXISTS idx_sessions_expires_at ON sessions(expires_at);
|
|
|
|
CREATE TABLE IF NOT EXISTS managed_domains (
|
|
name TEXT PRIMARY KEY COLLATE NOCASE,
|
|
added_at TEXT NOT NULL,
|
|
added_by INTEGER REFERENCES users(id)
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS shared_lists (
|
|
id INTEGER PRIMARY KEY CHECK (id = 1),
|
|
sender_denylist TEXT NOT NULL DEFAULT '[]',
|
|
sender_allowlist TEXT NOT NULL DEFAULT '[]',
|
|
recipient_denylist TEXT NOT NULL DEFAULT '[]'
|
|
);
|
|
|
|
INSERT OR IGNORE INTO shared_lists (id, sender_denylist, sender_allowlist, recipient_denylist)
|
|
VALUES (1, '[]', '[]', '[]');
|
|
`
|
|
_, err := s.db.Exec(schema)
|
|
return err
|
|
}
|
|
|
|
func (s *Store) CountUsers(ctx context.Context) (int, error) {
|
|
var count int
|
|
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&count)
|
|
return count, err
|
|
}
|
|
|
|
func (s *Store) CreateUser(ctx context.Context, email, displayName, passwordHash string, role Role) (*User, error) {
|
|
email = normalizeEmail(email)
|
|
displayName = strings.TrimSpace(displayName)
|
|
if email == "" {
|
|
return nil, ErrInvalidInput
|
|
}
|
|
if displayName == "" {
|
|
displayName = email
|
|
}
|
|
if role != RoleAdmin && role != RoleUser {
|
|
return nil, ErrInvalidInput
|
|
}
|
|
|
|
now := time.Now().UTC()
|
|
result, err := s.db.ExecContext(ctx, `
|
|
INSERT INTO users (email, display_name, role, password_hash, created_at)
|
|
VALUES (?, ?, ?, ?, ?)`,
|
|
email, displayName, string(role), passwordHash, now.Format(time.RFC3339Nano),
|
|
)
|
|
if err != nil {
|
|
if strings.Contains(strings.ToLower(err.Error()), "unique") {
|
|
return nil, ErrEmailTaken
|
|
}
|
|
return nil, err
|
|
}
|
|
id, err := result.LastInsertId()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &User{
|
|
ID: id,
|
|
Email: email,
|
|
DisplayName: displayName,
|
|
Role: role,
|
|
PasswordHash: passwordHash,
|
|
CreatedAt: now,
|
|
}, nil
|
|
}
|
|
|
|
func (s *Store) GetUserByID(ctx context.Context, id int64) (*User, error) {
|
|
row := s.db.QueryRowContext(ctx, `
|
|
SELECT id, email, display_name, role, password_hash, created_at
|
|
FROM users WHERE id = ?`, id)
|
|
return scanUser(row)
|
|
}
|
|
|
|
func (s *Store) GetUserByEmail(ctx context.Context, email string) (*User, error) {
|
|
row := s.db.QueryRowContext(ctx, `
|
|
SELECT id, email, display_name, role, password_hash, created_at
|
|
FROM users WHERE email = ? COLLATE NOCASE`, normalizeEmail(email))
|
|
return scanUser(row)
|
|
}
|
|
|
|
func (s *Store) ListUsers(ctx context.Context) ([]User, error) {
|
|
rows, err := s.db.QueryContext(ctx, `
|
|
SELECT id, email, display_name, role, password_hash, created_at
|
|
FROM users ORDER BY id ASC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
users := make([]User, 0)
|
|
for rows.Next() {
|
|
user, err := scanUser(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
users = append(users, *user)
|
|
}
|
|
return users, rows.Err()
|
|
}
|
|
|
|
func (s *Store) CreateSession(ctx context.Context, id string, userID int64, expiresAt time.Time) error {
|
|
now := time.Now().UTC()
|
|
_, err := s.db.ExecContext(ctx, `
|
|
INSERT INTO sessions (id, user_id, expires_at, created_at)
|
|
VALUES (?, ?, ?, ?)`,
|
|
id, userID, expiresAt.UTC().Format(time.RFC3339Nano), now.Format(time.RFC3339Nano),
|
|
)
|
|
return err
|
|
}
|
|
|
|
func (s *Store) GetSession(ctx context.Context, id string) (*Session, error) {
|
|
row := s.db.QueryRowContext(ctx, `
|
|
SELECT id, user_id, expires_at, created_at FROM sessions WHERE id = ?`, id)
|
|
|
|
var session Session
|
|
var expiresAt, createdAt string
|
|
if err := row.Scan(&session.ID, &session.UserID, &expiresAt, &createdAt); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, ErrNotFound
|
|
}
|
|
return nil, err
|
|
}
|
|
var err error
|
|
session.ExpiresAt, err = time.Parse(time.RFC3339Nano, expiresAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
session.CreatedAt, err = time.Parse(time.RFC3339Nano, createdAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &session, nil
|
|
}
|
|
|
|
func (s *Store) DeleteSession(ctx context.Context, id string) error {
|
|
_, err := s.db.ExecContext(ctx, `DELETE FROM sessions WHERE id = ?`, id)
|
|
return err
|
|
}
|
|
|
|
func (s *Store) DeleteExpiredSessions(ctx context.Context) error {
|
|
_, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM sessions WHERE expires_at < ?`, time.Now().UTC().Format(time.RFC3339Nano))
|
|
return err
|
|
}
|
|
|
|
func normalizeDomainName(name string) string {
|
|
return strings.ToLower(strings.TrimSpace(name))
|
|
}
|
|
|
|
func (s *Store) AddManagedDomain(ctx context.Context, name string, addedBy int64) (*ManagedDomain, error) {
|
|
name = normalizeDomainName(name)
|
|
if name == "" || !strings.Contains(name, ".") {
|
|
return nil, ErrInvalidInput
|
|
}
|
|
now := time.Now().UTC()
|
|
_, err := s.db.ExecContext(ctx, `
|
|
INSERT INTO managed_domains (name, added_at, added_by)
|
|
VALUES (?, ?, ?)`, name, now.Format(time.RFC3339Nano), addedBy)
|
|
if err != nil {
|
|
if strings.Contains(strings.ToLower(err.Error()), "unique") {
|
|
return nil, ErrDomainExists
|
|
}
|
|
return nil, err
|
|
}
|
|
return &ManagedDomain{
|
|
Name: name,
|
|
AddedAt: now,
|
|
AddedBy: &addedBy,
|
|
}, nil
|
|
}
|
|
|
|
func (s *Store) ListManagedDomains(ctx context.Context) ([]ManagedDomain, error) {
|
|
rows, err := s.db.QueryContext(ctx, `
|
|
SELECT name, added_at, added_by FROM managed_domains ORDER BY name ASC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
domains := make([]ManagedDomain, 0)
|
|
for rows.Next() {
|
|
var domain ManagedDomain
|
|
var addedAt string
|
|
var addedBy sql.NullInt64
|
|
if err := rows.Scan(&domain.Name, &addedAt, &addedBy); err != nil {
|
|
return nil, err
|
|
}
|
|
parsed, err := time.Parse(time.RFC3339Nano, addedAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
domain.AddedAt = parsed
|
|
if addedBy.Valid {
|
|
value := addedBy.Int64
|
|
domain.AddedBy = &value
|
|
}
|
|
domains = append(domains, domain)
|
|
}
|
|
return domains, rows.Err()
|
|
}
|
|
|
|
func (s *Store) GetManagedDomain(ctx context.Context, name string) (*ManagedDomain, error) {
|
|
row := s.db.QueryRowContext(ctx, `
|
|
SELECT name, added_at, added_by FROM managed_domains WHERE name = ? COLLATE NOCASE`,
|
|
normalizeDomainName(name))
|
|
|
|
var domain ManagedDomain
|
|
var addedAt string
|
|
var addedBy sql.NullInt64
|
|
if err := row.Scan(&domain.Name, &addedAt, &addedBy); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, ErrNotFound
|
|
}
|
|
return nil, err
|
|
}
|
|
parsed, err := time.Parse(time.RFC3339Nano, addedAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
domain.AddedAt = parsed
|
|
if addedBy.Valid {
|
|
value := addedBy.Int64
|
|
domain.AddedBy = &value
|
|
}
|
|
return &domain, nil
|
|
}
|
|
|
|
func (s *Store) IsManagedDomain(ctx context.Context, name string) (bool, error) {
|
|
_, err := s.GetManagedDomain(ctx, name)
|
|
if err == nil {
|
|
return true, nil
|
|
}
|
|
if errors.Is(err, ErrNotFound) {
|
|
return false, nil
|
|
}
|
|
return false, err
|
|
}
|
|
|
|
func (s *Store) DeleteManagedDomain(ctx context.Context, name string) error {
|
|
result, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM managed_domains WHERE name = ? COLLATE NOCASE`, normalizeDomainName(name))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
rows, err := result.RowsAffected()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if rows == 0 {
|
|
return ErrNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *Store) GetSharedLists(ctx context.Context) (*SharedLists, error) {
|
|
row := s.db.QueryRowContext(ctx, `
|
|
SELECT sender_denylist, sender_allowlist, recipient_denylist
|
|
FROM shared_lists WHERE id = 1`)
|
|
|
|
var denylistJSON, allowlistJSON, recipientJSON string
|
|
if err := row.Scan(&denylistJSON, &allowlistJSON, &recipientJSON); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return &SharedLists{
|
|
SenderDenylist: []string{},
|
|
SenderAllowlist: []string{},
|
|
RecipientDenylist: []string{},
|
|
}, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
|
|
lists := &SharedLists{}
|
|
if err := unmarshalStringList(denylistJSON, &lists.SenderDenylist); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := unmarshalStringList(allowlistJSON, &lists.SenderAllowlist); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := unmarshalStringList(recipientJSON, &lists.RecipientDenylist); err != nil {
|
|
return nil, err
|
|
}
|
|
return lists, nil
|
|
}
|
|
|
|
func (s *Store) PutSharedLists(ctx context.Context, lists SharedLists) error {
|
|
denylistJSON, err := marshalStringList(lists.SenderDenylist)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
allowlistJSON, err := marshalStringList(lists.SenderAllowlist)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
recipientJSON, err := marshalStringList(lists.RecipientDenylist)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
_, err = s.db.ExecContext(ctx, `
|
|
INSERT INTO shared_lists (id, sender_denylist, sender_allowlist, recipient_denylist)
|
|
VALUES (1, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET
|
|
sender_denylist = excluded.sender_denylist,
|
|
sender_allowlist = excluded.sender_allowlist,
|
|
recipient_denylist = excluded.recipient_denylist`,
|
|
denylistJSON, allowlistJSON, recipientJSON,
|
|
)
|
|
return err
|
|
}
|
|
|
|
func marshalStringList(values []string) (string, error) {
|
|
normalized := normalizeStringList(values)
|
|
data, err := json.Marshal(normalized)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return string(data), nil
|
|
}
|
|
|
|
func unmarshalStringList(raw string, dest *[]string) error {
|
|
if strings.TrimSpace(raw) == "" {
|
|
*dest = []string{}
|
|
return nil
|
|
}
|
|
var values []string
|
|
if err := json.Unmarshal([]byte(raw), &values); err != nil {
|
|
return err
|
|
}
|
|
*dest = normalizeStringList(values)
|
|
return nil
|
|
}
|
|
|
|
func normalizeStringList(values []string) []string {
|
|
return lists.Compact(values)
|
|
}
|
|
|
|
type scannable interface {
|
|
Scan(dest ...any) error
|
|
}
|
|
|
|
func scanUser(row scannable) (*User, error) {
|
|
var user User
|
|
var role string
|
|
var createdAt string
|
|
if err := row.Scan(&user.ID, &user.Email, &user.DisplayName, &role, &user.PasswordHash, &createdAt); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, ErrNotFound
|
|
}
|
|
return nil, err
|
|
}
|
|
user.Role = Role(role)
|
|
parsed, err := time.Parse(time.RFC3339Nano, createdAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
user.CreatedAt = parsed
|
|
return &user, nil
|
|
}
|
|
|
|
func normalizeEmail(email string) string {
|
|
return strings.ToLower(strings.TrimSpace(email))
|
|
}
|