From 38475e4b61ab79e3d4d4eb2e8892b6c6a9a59045 Mon Sep 17 00:00:00 2001 From: Lewis Date: Tue, 19 May 2026 13:13:29 +0300 Subject: [PATCH] rbac,knotserver,spindle: txn + querier for member acl Lewis: May this revision serve well! --- knotserver/db/db.go | 28 ++++++++++ knotserver/db/known_dids.go | 8 +-- knotserver/db/member.go | 99 +++++++++++++++++++++++++++++++++ knotserver/router.go | 4 +- rbac/rbac.go | 24 +++++++- rbac/txn.go | 41 ++++++++++++++ rbac/util.go | 25 +++++++-- spindle/db/db.go | 108 +++++++++++++++++++++++++++++++++++- spindle/db/known_dids.go | 8 +-- spindle/db/member.go | 30 +++++----- 10 files changed, 342 insertions(+), 33 deletions(-) create mode 100644 knotserver/db/member.go create mode 100644 rbac/txn.go diff --git a/knotserver/db/db.go b/knotserver/db/db.go index f655fb03..00674ed0 100644 --- a/knotserver/db/db.go +++ b/knotserver/db/db.go @@ -24,6 +24,18 @@ type Querier interface { Exec(query string, args ...any) (sql.Result, error) } +func (d *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error) { + return d.db.BeginTx(ctx, opts) +} + +func (d *DB) Exec(query string, args ...any) (sql.Result, error) { + return d.db.Exec(query, args...) +} + +func (d *DB) QueryRow(query string, args ...any) *sql.Row { + return d.db.QueryRow(query, args...) +} + func Setup(ctx context.Context, dbPath string) (*DB, error) { // https://github.com/mattn/go-sqlite3#connection-string opts := []string{ @@ -179,6 +191,22 @@ func Setup(ctx context.Context, dbPath string) (*DB, error) { return nil, err } + if err := orm.RunMigration(conn, logger, "create-knot-members", func(tx *sql.Tx) error { + _, mErr := tx.ExecContext(ctx, ` + create table if not exists knot_members ( + id integer primary key autoincrement, + did text not null, + rkey text not null, + subject text not null, + created text not null default (strftime('%Y-%m-%dT%H:%M:%SZ', 'now')), + unique (did, rkey) + ); + `) + return mErr + }); err != nil { + return nil, err + } + return &DB{ db: db, logger: logger, diff --git a/knotserver/db/known_dids.go b/knotserver/db/known_dids.go index e010efb1..9d575250 100644 --- a/knotserver/db/known_dids.go +++ b/knotserver/db/known_dids.go @@ -1,12 +1,12 @@ package db -func (d *DB) AddDid(did string) error { - _, err := d.db.Exec(`insert or ignore into known_dids (did) values (?)`, did) +func AddDid(q DBTX, did string) error { + _, err := q.Exec(`insert or ignore into known_dids (did) values (?)`, did) return err } -func (d *DB) RemoveDid(did string) error { - _, err := d.db.Exec(`delete from known_dids where did = ?`, did) +func RemoveDid(q DBTX, did string) error { + _, err := q.Exec(`delete from known_dids where did = ?`, did) return err } diff --git a/knotserver/db/member.go b/knotserver/db/member.go new file mode 100644 index 00000000..9bb06f90 --- /dev/null +++ b/knotserver/db/member.go @@ -0,0 +1,99 @@ +package db + +import ( + "context" + "database/sql" + + "github.com/bluesky-social/indigo/atproto/syntax" + "tangled.org/core/orm" +) + +type KnotMember struct { + Id int + Did syntax.DID + Rkey string + Subject syntax.DID +} + +func (d *DB) IsMigrationApplied(name string) (bool, error) { + var exists bool + err := d.db.QueryRow( + `select exists (select 1 from migrations where name = ?)`, + name, + ).Scan(&exists) + return exists, err +} + +func (d *DB) ApplyKnotMembersBackfill(ctx context.Context, rows []KnotMember, migrationName string) error { + conn, err := d.db.Conn(ctx) + if err != nil { + return err + } + defer conn.Close() + + return orm.RunMigration(conn, d.logger, migrationName, func(tx *sql.Tx) error { + for _, m := range rows { + if _, err := tx.ExecContext(ctx, + `insert or ignore into known_dids (did) values (?)`, + m.Subject, + ); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, + `insert or ignore into knot_members (did, rkey, subject) values (?, ?, ?)`, + m.Did, m.Rkey, m.Subject, + ); err != nil { + return err + } + } + return nil + }) +} + +func AddKnotMember(q DBTX, member KnotMember) error { + _, err := q.Exec( + `insert or ignore into knot_members (did, rkey, subject) values (?, ?, ?)`, + member.Did, + member.Rkey, + member.Subject, + ) + return err +} + +func RemoveKnotMember(q DBTX, ownerDid, rkey string) error { + _, err := q.Exec( + "delete from knot_members where did = ? and rkey = ?", + ownerDid, + rkey, + ) + return err +} + +func CountKnotMembersBySubject(q DBTX, subject string) (int, error) { + var count int + err := q.QueryRow( + `select count(*) from knot_members where subject = ?`, + subject, + ).Scan(&count) + return count, err +} + +func GetKnotMember(q DBTX, did, rkey string) (*KnotMember, error) { + query := + `select id, did, rkey, subject + from knot_members + where did = ? and rkey = ?` + + var member KnotMember + err := q.QueryRow(query, did, rkey).Scan( + &member.Id, + &member.Did, + &member.Rkey, + &member.Subject, + ) + if err != nil { + return nil, err + } + + return &member, nil +} diff --git a/knotserver/router.go b/knotserver/router.go index d742623a..6e4b366c 100644 --- a/knotserver/router.go +++ b/knotserver/router.go @@ -194,7 +194,7 @@ func (h *Knot) configureOwner(ctx context.Context) error { } // remove existing owner - if err = h.db.RemoveDid(existingOwner); err != nil { + if err = db.RemoveDid(h.db, existingOwner); err != nil { return err } if err = h.e.RemoveKnotOwner(rbacDomain, existingOwner); err != nil { @@ -205,7 +205,7 @@ func (h *Knot) configureOwner(ctx context.Context) error { return fmt.Errorf("more than one owner in DB, try deleting %q and starting over", h.c.Server.DBPath) } - if err = h.db.AddDid(cfgOwner); err != nil { + if err = db.AddDid(h.db, cfgOwner); err != nil { return fmt.Errorf("failed to add owner to DB: %w", err) } if err := h.e.AddKnotOwner(rbacDomain, cfgOwner); err != nil { diff --git a/rbac/rbac.go b/rbac/rbac.go index 42b3ef31..6887a648 100644 --- a/rbac/rbac.go +++ b/rbac/rbac.go @@ -126,10 +126,20 @@ func (e *Enforcer) RemoveKnotOwner(domain, owner string) error { } func (e *Enforcer) AddKnotMember(domain, member string) error { - return e.addMember(domain, member) + _, err := e.addMember(domain, member) + return err } func (e *Enforcer) RemoveKnotMember(domain, member string) error { + _, err := e.removeMember(domain, member) + return err +} + +func (e *Enforcer) TryAddKnotMember(domain, member string) (bool, error) { + return e.addMember(domain, member) +} + +func (e *Enforcer) TryRemoveKnotMember(domain, member string) (bool, error) { return e.removeMember(domain, member) } @@ -142,10 +152,20 @@ func (e *Enforcer) RemoveSpindleOwner(domain, owner string) error { } func (e *Enforcer) AddSpindleMember(domain, member string) error { - return e.addMember(intoSpindle(domain), member) + _, err := e.addMember(intoSpindle(domain), member) + return err } func (e *Enforcer) RemoveSpindleMember(domain, member string) error { + _, err := e.removeMember(intoSpindle(domain), member) + return err +} + +func (e *Enforcer) TryAddSpindleMember(domain, member string) (bool, error) { + return e.addMember(intoSpindle(domain), member) +} + +func (e *Enforcer) TryRemoveSpindleMember(domain, member string) (bool, error) { return e.removeMember(intoSpindle(domain), member) } diff --git a/rbac/txn.go b/rbac/txn.go new file mode 100644 index 00000000..89a33f38 --- /dev/null +++ b/rbac/txn.go @@ -0,0 +1,41 @@ +package rbac + +import ( + "database/sql" + "log/slog" + "slices" +) + +type Txn struct { + SQLTx *sql.Tx + undos []func() error + committed bool +} + +func NewTxn(sqlTx *sql.Tx) *Txn { + return &Txn{SQLTx: sqlTx} +} + +func (t *Txn) AddUndo(undo func() error) { + t.undos = append(t.undos, undo) +} + +func (t *Txn) Commit() error { + if err := t.SQLTx.Commit(); err != nil { + return err + } + t.committed = true + return nil +} + +func (t *Txn) Cleanup(l *slog.Logger) { + if t.committed { + return + } + t.SQLTx.Rollback() + for _, undo := range slices.Backward(t.undos) { + if err := undo(); err != nil { + l.Error("failed to reverse ACL change", "err", err) + } + } +} diff --git a/rbac/util.go b/rbac/util.go index e67c3d4b..41ad2495 100644 --- a/rbac/util.go +++ b/rbac/util.go @@ -34,14 +34,12 @@ func (e *Enforcer) removeOwner(domain, owner string) error { return err } -func (e *Enforcer) addMember(domain, member string) error { - _, err := e.E.AddGroupingPolicy(member, "server:member", domain) - return err +func (e *Enforcer) addMember(domain, member string) (bool, error) { + return e.E.AddGroupingPolicy(member, "server:member", domain) } -func (e *Enforcer) removeMember(domain, member string) error { - _, err := e.E.RemoveGroupingPolicy(member, "server:member", domain) - return err +func (e *Enforcer) removeMember(domain, member string) (bool, error) { + return e.E.RemoveGroupingPolicy(member, "server:member", domain) } func (e *Enforcer) isRole(user, role, domain string) (bool, error) { @@ -59,6 +57,21 @@ func (e *Enforcer) isInviteAllowed(user, domain string) (bool, error) { return e.E.Enforce(user, domain, domain, "server:invite") } +func (e *Enforcer) HasAnyPolicyForUser(user string) (bool, error) { + pPolicies, err := e.E.GetFilteredNamedPolicy("p", 0, user) + if err != nil { + return false, err + } + if len(pPolicies) > 0 { + return true, nil + } + gPolicies, err := e.E.GetFilteredNamedGroupingPolicy("g", 0, user) + if err != nil { + return false, err + } + return len(gPolicies) > 0, nil +} + func checkRepoFormat(repo string) error { // sanity check, repo must be of the form ownerDid/repo if parts := strings.SplitN(repo, "/", 2); !strings.HasPrefix(parts[0], "did:") { diff --git a/spindle/db/db.go b/spindle/db/db.go index 0d89f357..57d4dfb1 100644 --- a/spindle/db/db.go +++ b/spindle/db/db.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "log/slog" + "slices" "strings" _ "github.com/mattn/go-sqlite3" @@ -15,6 +16,11 @@ type DB struct { *sql.DB } +type DBTX interface { + QueryRow(query string, args ...any) *sql.Row + Exec(query string, args ...any) (sql.Result, error) +} + func Make(ctx context.Context, dbPath string) (*DB, error) { // https://github.com/mattn/go-sqlite3#connection-string opts := []string{ @@ -84,7 +90,7 @@ func Make(ctx context.Context, dbPath string) (*DB, error) { created text not null default (strftime('%Y-%m-%dT%H:%M:%SZ', 'now')), -- constraints - unique (did, instance, subject) + unique (did, rkey) ); -- status event for a single workflow @@ -112,7 +118,7 @@ func Make(ctx context.Context, dbPath string) (*DB, error) { } func runMigrations(_ context.Context, conn *sql.Conn, logger *slog.Logger) error { - return orm.RunMigration(conn, logger, "repos-to-repo-did", func(tx *sql.Tx) error { + if err := orm.RunMigration(conn, logger, "repos-to-repo-did", func(tx *sql.Tx) error { var hasName int if err := tx.QueryRow( `select count(*) from pragma_table_info('repos') where name = 'name'`, @@ -162,9 +168,107 @@ func runMigrations(_ context.Context, conn *sql.Conn, logger *slog.Logger) error on repo_collaborators(repo_did); `) return err + }); err != nil { + return err + } + + return orm.RunMigration(conn, logger, "spindle-members-unique-on-rkey", func(tx *sql.Tx) error { + hasTarget, err := hasUniqueIndex(tx, "spindle_members", []string{"did", "rkey"}) + if err != nil { + return err + } + if hasTarget { + return nil + } + + var totalRows, distinctRows int + if err := tx.QueryRow(`select count(*) from spindle_members`).Scan(&totalRows); err != nil { + return err + } + if err := tx.QueryRow(`select count(*) from (select 1 from spindle_members group by did, rkey)`).Scan(&distinctRows); err != nil { + return err + } + if dropped := totalRows - distinctRows; dropped > 0 { + logger.Warn("dropping duplicate (did, rkey) rows during spindle_members rebuild", "dropped", dropped, "kept", distinctRows) + } + + _, err = tx.Exec(` + create table spindle_members_new ( + id integer primary key autoincrement, + did text not null, + rkey text not null, + instance text not null, + subject text not null, + created text not null default (strftime('%Y-%m-%dT%H:%M:%SZ', 'now')), + unique (did, rkey) + ); + + insert into spindle_members_new (id, did, rkey, instance, subject, created) + select id, did, rkey, instance, subject, created + from spindle_members sm + where id = ( + select max(id) from spindle_members + where did = sm.did and rkey = sm.rkey + ); + + drop table spindle_members; + alter table spindle_members_new rename to spindle_members; + `) + return err }) } +func hasUniqueIndex(tx *sql.Tx, table string, cols []string) (bool, error) { + rows, err := tx.Query( + `select name from pragma_index_list(?) where "unique" = 1`, + table, + ) + if err != nil { + return false, err + } + defer rows.Close() + + var indexNames []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return false, err + } + indexNames = append(indexNames, name) + } + if err := rows.Err(); err != nil { + return false, err + } + + wantSorted := slices.Clone(cols) + slices.Sort(wantSorted) + + for _, name := range indexNames { + colRows, err := tx.Query( + `select name from pragma_index_info(?) order by seqno`, + name, + ) + if err != nil { + return false, err + } + var got []string + for colRows.Next() { + var c string + if err := colRows.Scan(&c); err != nil { + colRows.Close() + return false, err + } + got = append(got, c) + } + colRows.Close() + slices.Sort(got) + if slices.Equal(got, wantSorted) { + return true, nil + } + } + return false, nil +} + func (d *DB) SaveLastTimeUs(lastTimeUs int64) error { _, err := d.Exec(` insert into _jetstream (id, last_time_us) diff --git a/spindle/db/known_dids.go b/spindle/db/known_dids.go index 8acdb695..2192e069 100644 --- a/spindle/db/known_dids.go +++ b/spindle/db/known_dids.go @@ -1,12 +1,12 @@ package db -func (d *DB) AddDid(did string) error { - _, err := d.Exec(`insert or ignore into known_dids (did) values (?)`, did) +func AddDid(q DBTX, did string) error { + _, err := q.Exec(`insert or ignore into known_dids (did) values (?)`, did) return err } -func (d *DB) RemoveDid(did string) error { - _, err := d.Exec(`delete from known_dids where did = ?`, did) +func RemoveDid(q DBTX, did string) error { + _, err := q.Exec(`delete from known_dids where did = ?`, did) return err } diff --git a/spindle/db/member.go b/spindle/db/member.go index a93b7f91..31fbfe19 100644 --- a/spindle/db/member.go +++ b/spindle/db/member.go @@ -1,8 +1,6 @@ package db import ( - "time" - "github.com/bluesky-social/indigo/atproto/syntax" ) @@ -12,11 +10,10 @@ type SpindleMember struct { Rkey string // rkey of the record Instance string Subject syntax.DID // the member being added - Created time.Time } -func AddSpindleMember(db *DB, member SpindleMember) error { - _, err := db.Exec( +func AddSpindleMember(q DBTX, member SpindleMember) error { + _, err := q.Exec( `insert or ignore into spindle_members (did, rkey, instance, subject) values (?, ?, ?, ?)`, member.Did, member.Rkey, @@ -26,30 +23,37 @@ func AddSpindleMember(db *DB, member SpindleMember) error { return err } -func RemoveSpindleMember(db *DB, owner_did, rkey string) error { - _, err := db.Exec( +func RemoveSpindleMember(q DBTX, ownerDid, rkey string) error { + _, err := q.Exec( "delete from spindle_members where did = ? and rkey = ?", - owner_did, + ownerDid, rkey, ) return err } -func GetSpindleMember(db *DB, did, rkey string) (*SpindleMember, error) { +func CountSpindleMembersBySubject(q DBTX, subject string) (int, error) { + var count int + err := q.QueryRow( + `select count(*) from spindle_members where subject = ?`, + subject, + ).Scan(&count) + return count, err +} + +func GetSpindleMember(q DBTX, did, rkey string) (*SpindleMember, error) { query := - `select id, did, rkey, instance, subject, created + `select id, did, rkey, instance, subject from spindle_members where did = ? and rkey = ?` var member SpindleMember - var createdAt string - err := db.QueryRow(query, did, rkey).Scan( + err := q.QueryRow(query, did, rkey).Scan( &member.Id, &member.Did, &member.Rkey, &member.Instance, &member.Subject, - &createdAt, ) if err != nil { return nil, err -- 2.51.2