diff --git a/internal/sites/sites.go b/internal/sites/sites.go index 7d274eb..a785172 100644 --- a/internal/sites/sites.go +++ b/internal/sites/sites.go @@ -72,10 +72,19 @@ func (m *Manager) Add(domain, username string) (string, error) { if !Valid(domain) { return "", ErrInvalidDomain } - if err := m.st.MapDomain(domain, username); err != nil { + created, err := m.st.MapDomain(domain, username) + if err != nil { return domain, err } - return domain, m.link(domain, username) + if err := m.link(domain, username); err != nil { + if created { + if rollbackErr := m.st.UnmapDomain(domain, username); rollbackErr != nil { + return domain, errors.Join(err, rollbackErr) + } + } + return domain, err + } + return domain, nil } // Remove unmaps a domain owned by username (DB row + symlink). diff --git a/internal/sites/sites_test.go b/internal/sites/sites_test.go index 7108520..4d2fec0 100644 --- a/internal/sites/sites_test.go +++ b/internal/sites/sites_test.go @@ -111,6 +111,73 @@ func TestManagerAddRemoveSyncAsk(t *testing.T) { } } +func TestManagerAddRollsBackNewMappingWhenLinkCreationFails(t *testing.T) { + dir := t.TempDir() + st, err := store.Open(filepath.Join(dir, "test.db")) + if err != nil { + t.Fatal(err) + } + defer st.Close() + + m, err := NewManager(st, dir) + if err != nil { + t.Fatal(err) + } + domain := "blocked.example.com" + blocked := filepath.Join(dir, "domains", domain) + if err := os.MkdirAll(blocked, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(blocked, "keep"), []byte("occupied"), 0o644); err != nil { + t.Fatal(err) + } + + if _, err := m.Add(domain, "alice"); err == nil { + t.Fatal("Add succeeded despite an occupied domain link path") + } + if owner, ok, err := st.DomainUser(domain); err != nil { + t.Fatal(err) + } else if ok { + t.Fatalf("failed Add left domain mapped to %q", owner) + } +} + +func TestManagerAddPreservesExistingMappingWhenLinkRepairFails(t *testing.T) { + dir := t.TempDir() + st, err := store.Open(filepath.Join(dir, "test.db")) + if err != nil { + t.Fatal(err) + } + defer st.Close() + + m, err := NewManager(st, dir) + if err != nil { + t.Fatal(err) + } + domain := "blocked.example.com" + if created, err := st.MapDomain(domain, "alice"); err != nil { + t.Fatal(err) + } else if !created { + t.Fatal("expected a new domain mapping") + } + blocked := filepath.Join(dir, "domains", domain) + if err := os.MkdirAll(blocked, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(blocked, "keep"), []byte("occupied"), 0o644); err != nil { + t.Fatal(err) + } + + if _, err := m.Add(domain, "alice"); err == nil { + t.Fatal("Add succeeded despite an occupied domain link path") + } + if owner, ok, err := st.DomainUser(domain); err != nil { + t.Fatal(err) + } else if !ok || owner != "alice" { + t.Fatalf("existing mapping changed after failed repair: owner=%q ok=%v", owner, ok) + } +} + func TestSyncRemovesStaleDomainEntries(t *testing.T) { dir := t.TempDir() st, err := store.Open(filepath.Join(dir, "test.db")) diff --git a/internal/store/store.go b/internal/store/store.go index f88cf64..12fec43 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -136,9 +136,10 @@ type Store interface { OnlineUsers() (map[string]bool, error) // Custom domains mapped to a member's homepage (public_html). - // MapDomain binds domain→username, returning ErrDomainTaken if it is - // already claimed by someone else (re-binding to the same owner is a no-op). - MapDomain(domain, username string) error + // MapDomain binds domain→username, returning whether a new row was created. + // It returns ErrDomainTaken if the domain is already claimed by someone else; + // re-binding to the same owner is a no-op and returns false. + MapDomain(domain, username string) (bool, error) UnmapDomain(domain, username string) error DomainUser(domain string) (string, bool, error) DomainsForUser(username string) ([]string, error) @@ -870,20 +871,20 @@ func (s *sqliteStore) OnlineUsers() (map[string]bool, error) { return online, rows.Err() } -func (s *sqliteStore) MapDomain(domain, username string) error { +func (s *sqliteStore) MapDomain(domain, username string) (bool, error) { var owner string err := s.db.QueryRow(`SELECT username FROM domains WHERE domain = ?`, domain).Scan(&owner) switch { case err == nil: if owner != username { - return ErrDomainTaken + return false, ErrDomainTaken } - return nil // already ours + return false, nil // already ours case err != sql.ErrNoRows: - return err + return false, err } _, err = s.db.Exec(`INSERT INTO domains (domain, username) VALUES (?,?)`, domain, username) - return err + return err == nil, err } func (s *sqliteStore) UnmapDomain(domain, username string) error {