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)) }