fix(sites): roll back failed domain claims (#119)
Some checks are pending
CI / build (push) Waiting to run
deploy / deploy (push) Waiting to run
test / test (push) Waiting to run

* test(sites): reproduce stale failed domain claims

Signed-off-by: RissRIce <jsdavid278@gmail.com>

* fix(sites): roll back failed domain claims

Signed-off-by: RissRIce <jsdavid278@gmail.com>

---------

Signed-off-by: RissRIce <jsdavid278@gmail.com>
This commit is contained in:
RissRIce 2026-08-13 03:34:30 -06:00 committed by GitHub
parent 5a1f5d90db
commit 4b815eaab8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 87 additions and 10 deletions

View file

@ -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).

View file

@ -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"))

View file

@ -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 {