diff --git a/knotserver/db/db.go b/knotserver/db/db.go index b7fdbec5..f655fb03 100644 --- a/knotserver/db/db.go +++ b/knotserver/db/db.go @@ -193,38 +193,48 @@ func (d *DB) StoreRepoDidWeb(repoDid, ownerDid, repoName string) error { return d.storeRepoKeyRow(repoDid, nil, ownerDid, repoName, "web") } -func (d *DB) storeRepoKeyRow(repoDid string, signingKey []byte, ownerDid, repoName, keyType string) (err error) { +func (d *DB) storeRepoKeyRow(repoDid string, signingKey []byte, ownerDid, repoName, keyType string) error { tx, err := d.db.Begin() if err != nil { return err } - defer func() { - if err != nil { - tx.Rollback() - return - } - err = tx.Commit() - }() - - if _, err = tx.Exec( + defer tx.Rollback() + + if _, err := tx.Exec( `INSERT INTO repo_keys (repo_did, signing_key, owner_did, repo_name, key_type) VALUES (?, ?, ?, ?, ?)`, repoDid, signingKey, ownerDid, repoName, keyType, ); err != nil { return err } - _, err = tx.Exec( + if _, err := tx.Exec( `INSERT INTO repo_aliases (owner_did, rkey, repo_did, rev) VALUES (?, ?, ?, '0_' || strftime('%Y-%m-%dT%H:%M:%SZ', 'now')) ON CONFLICT(owner_did, rkey) DO NOTHING`, ownerDid, repoName, repoDid, - ) - return err + ); err != nil { + return err + } + + return tx.Commit() } func (d *DB) DeleteRepoKey(repoDid string) error { - _, err := d.db.Exec(`DELETE FROM repo_keys WHERE repo_did = ?`, repoDid) - return err + tx, err := d.db.Begin() + if err != nil { + return err + } + defer tx.Rollback() + + if _, err := tx.Exec(`DELETE FROM repo_aliases WHERE repo_did = ?`, repoDid); err != nil { + return err + } + + if _, err := tx.Exec(`DELETE FROM repo_keys WHERE repo_did = ?`, repoDid); err != nil { + return err + } + + return tx.Commit() } func (d *DB) RepoDidExists(repoDid string) (bool, error) { @@ -242,6 +252,15 @@ func (d *DB) GetRepoDid(ownerDid, rkey string) (string, error) { return repoDid, err } +func (d *DB) GetRepoDidByName(ownerDid, repoName string) (string, error) { + var repoDid string + err := d.db.QueryRow( + `SELECT repo_did FROM repo_keys WHERE owner_did = ? AND repo_name = ?`, + ownerDid, repoName, + ).Scan(&repoDid) + return repoDid, err +} + func (d *DB) GetRepoKeyOwner(repoDid string) (string, string, error) { return GetRepoKeyOwner(d.db, repoDid) } diff --git a/knotserver/db/repo_aliases.go b/knotserver/db/repo_aliases.go index 7343cd88..e497635a 100644 --- a/knotserver/db/repo_aliases.go +++ b/knotserver/db/repo_aliases.go @@ -25,14 +25,6 @@ func (d *DB) UpsertRepoAlias(a RepoAlias) error { return err } -func (d *DB) DeleteRepoAlias(ownerDid, rkey string) error { - _, err := d.db.Exec( - `delete from repo_aliases where owner_did = ? and rkey = ?`, - ownerDid, rkey, - ) - return err -} - func (d *DB) ResolveAlias(ownerDid, rkey string) (*RepoAlias, error) { var a RepoAlias err := d.db.QueryRow( diff --git a/knotserver/ingester.go b/knotserver/ingester.go index bba0e082..de491240 100644 --- a/knotserver/ingester.go +++ b/knotserver/ingester.go @@ -499,9 +499,6 @@ func (h *Knot) processRepo(ctx context.Context, event *jmodels.Event) error { } if event.Commit.Operation == jmodels.CommitOperationDelete { - if err := h.db.DeleteRepoAlias(event.Did, rkey); err != nil { - l.Warn("failed to delete repo alias", "err", err) - } return nil } diff --git a/knotserver/ingester_repo_test.go b/knotserver/ingester_repo_test.go index 54c2a8d7..4b14c2d8 100644 --- a/knotserver/ingester_repo_test.go +++ b/knotserver/ingester_repo_test.go @@ -4,6 +4,8 @@ import ( "context" "encoding/json" "log/slog" + "os" + "path/filepath" "sync" "testing" @@ -12,6 +14,7 @@ import ( "tangled.org/core/knotserver/config" "tangled.org/core/knotserver/db" "tangled.org/core/log" + "tangled.org/core/rbac" ) type logRecord struct { @@ -71,17 +74,33 @@ func (h *capturingHandler) snapshot() []logRecord { func newProcessRepoFixture(t *testing.T) (*Knot, context.Context, *capturingHandler) { t.Helper() - d := newTestKnotDB(t) + scanPath := t.TempDir() + dbPath := filepath.Join(scanPath, "knot.db") + d, err := db.Setup(context.Background(), dbPath) + if err != nil { + t.Fatalf("db.Setup: %v", err) + } + + e, err := rbac.NewEnforcer(dbPath) + if err != nil { + t.Fatalf("rbac.NewEnforcer: %v", err) + } + if err := e.AddKnot(rbac.ThisServer); err != nil { + t.Fatalf("AddKnot: %v", err) + } + cap := newCapturingHandler() l := slog.New(cap) ctx := log.IntoContext(context.Background(), l) c := &config.Config{ Server: config.Server{Hostname: "knot.example"}, + Repo: config.Repo{ScanPath: scanPath}, } return &Knot{ c: c, db: d, + e: e, l: l, }, ctx, cap } @@ -135,7 +154,7 @@ func TestProcessRepo_CreateRegistersAlias(t *testing.T) { } } -func TestProcessRepo_DeleteRemovesAlias(t *testing.T) { +func TestProcessRepo_DeleteIsNoOp(t *testing.T) { h, ctx, _ := newProcessRepoFixture(t) if err := h.db.StoreRepoKey("did:plc:repo1", []byte("k"), "did:plc:akshay", "foo"); err != nil { t.Fatalf("StoreRepoKey: %v", err) @@ -145,19 +164,33 @@ func TestProcessRepo_DeleteRemovesAlias(t *testing.T) { }); err != nil { t.Fatalf("UpsertRepoAlias: %v", err) } + if err := h.e.AddRepo("did:plc:akshay", rbac.ThisServer, "did:plc:repo1"); err != nil { + t.Fatalf("AddRepo rbac: %v", err) + } + repoPath := filepath.Join(h.c.Repo.ScanPath, "did:plc:repo1") + if err := os.MkdirAll(repoPath, 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } ev := repoEvent(t, "did:plc:akshay", "bar", "3laaaaaaaaaac", tangled.Repo{}, jsmodels.CommitOperationDelete) if err := h.processRepo(ctx, ev); err != nil { t.Fatalf("processRepo: %v", err) } - if _, err := h.db.GetRepoDid("did:plc:akshay", "bar"); err == nil { - t.Errorf("bar alias should have been deleted") + if got, err := h.db.GetRepoDid("did:plc:akshay", "bar"); err != nil || got != "did:plc:repo1" { + t.Errorf("bar alias should be untouched by firehose delete: got (%q, %v)", got, err) } - - _, current, _ := h.db.CurrentRkey("did:plc:repo1") - if current != "foo" { - t.Errorf("current rkey after delete = %q, want foo", current) + if got, err := h.db.GetRepoDid("did:plc:akshay", "foo"); err != nil || got != "did:plc:repo1" { + t.Errorf("foo alias should be untouched by firehose delete: got (%q, %v)", got, err) + } + if exists, _ := h.db.RepoDidExists("did:plc:repo1"); !exists { + t.Errorf("repo_keys row should be untouched by firehose delete") + } + if _, err := os.Stat(repoPath); err != nil { + t.Errorf("repo dir should be untouched by firehose delete: %v", err) + } + if allowed, _ := h.e.IsRepoDeleteAllowed("did:plc:akshay", rbac.ThisServer, "did:plc:repo1"); !allowed { + t.Errorf("rbac policies should be untouched by firehose delete") } } diff --git a/knotserver/xrpc/create_repo.go b/knotserver/xrpc/create_repo.go index 63fbd281..5dffc561 100644 --- a/knotserver/xrpc/create_repo.go +++ b/knotserver/xrpc/create_repo.go @@ -2,6 +2,7 @@ package xrpc import ( "context" + "database/sql" "encoding/json" "errors" "fmt" @@ -103,7 +104,23 @@ func (h *Xrpc) CreateRepo(w http.ResponseWriter, r *http.Request) { return default: + removeOrphan := func(orphanDid string) error { + orphanPath, _ := securejoin.SecureJoin(h.Config.Repo.ScanPath, orphanDid) + if rmErr := os.RemoveAll(orphanPath); rmErr != nil { + l.Warn("failed to remove orphan repo directory", "path", orphanPath, "error", rmErr.Error()) + } + if rbacErr := h.Enforcer.RemoveRepo(actorDid.String(), rbac.ThisServer, orphanDid); rbacErr != nil { + l.Warn("failed to remove orphan rbac entry", "repoDid", orphanDid, "error", rbacErr.Error()) + } + return h.Db.DeleteRepoKey(orphanDid) + } + existingDid, dbErr := h.Db.GetRepoDid(actorDid.String(), repoName) + if dbErr != nil && !errors.Is(dbErr, sql.ErrNoRows) { + l.Error("failed to look up repo alias", "error", dbErr.Error()) + writeError(w, xrpcerr.GenericError(dbErr), http.StatusInternalServerError) + return + } if dbErr == nil && existingDid != "" { didRepoPath, _ := securejoin.SecureJoin(h.Config.Repo.ScanPath, existingDid) if _, statErr := os.Stat(didRepoPath); statErr == nil { @@ -113,11 +130,32 @@ func (h *Xrpc) CreateRepo(w http.ResponseWriter, r *http.Request) { return } l.Warn("stale repo key found without directory, cleaning up", "repoDid", existingDid) - if delErr := h.Db.DeleteRepoKey(existingDid); delErr != nil { + if delErr := removeOrphan(existingDid); delErr != nil { l.Error("failed to clean up stale repo key", "repoDid", existingDid, "error", delErr.Error()) writeError(w, xrpcerr.GenericError(fmt.Errorf("failed to clean up stale state, retry later")), http.StatusInternalServerError) return } + } else { + orphanDid, lookupErr := h.Db.GetRepoDidByName(actorDid.String(), repoName) + if lookupErr != nil && !errors.Is(lookupErr, sql.ErrNoRows) { + l.Error("failed to look up orphan repo key", "error", lookupErr.Error()) + writeError(w, xrpcerr.GenericError(lookupErr), http.StatusInternalServerError) + return + } + if lookupErr == nil && orphanDid != "" { + orphanPath, _ := securejoin.SecureJoin(h.Config.Repo.ScanPath, orphanDid) + if _, statErr := os.Stat(orphanPath); statErr == nil { + l.Error("orphan repo_keys row but directory present, refusing to overwrite", "repoDid", orphanDid) + writeError(w, xrpcerr.GenericError(fmt.Errorf("repository %q is in an inconsistent state, contact a knot admin", repoName)), http.StatusConflict) + return + } + l.Warn("orphan repo_keys row without alias, cleaning up", "repoDid", orphanDid) + if delErr := removeOrphan(orphanDid); delErr != nil { + l.Error("failed to clean up orphan repo key", "repoDid", orphanDid, "error", delErr.Error()) + writeError(w, xrpcerr.GenericError(fmt.Errorf("failed to clean up orphan state, retry later")), http.StatusInternalServerError) + return + } + } } var prepErr error diff --git a/knotserver/xrpc/delete_repo.go b/knotserver/xrpc/delete_repo.go index 376b80cf..4df09d6f 100644 --- a/knotserver/xrpc/delete_repo.go +++ b/knotserver/xrpc/delete_repo.go @@ -1,7 +1,9 @@ package xrpc import ( + "database/sql" "encoding/json" + "errors" "fmt" "net/http" "os" @@ -9,6 +11,7 @@ import ( comatproto "github.com/bluesky-social/indigo/api/atproto" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/bluesky-social/indigo/xrpc" + securejoin "github.com/cyphar/filepath-securejoin" "tangled.org/core/api/tangled" "tangled.org/core/rbac" xrpcerr "tangled.org/core/xrpc/errors" @@ -60,13 +63,23 @@ func (x *Xrpc) DeleteRepo(w http.ResponseWriter, r *http.Request) { } repoDid, err := x.Db.GetRepoDid(did, name) + if errors.Is(err, sql.ErrNoRows) { + repoDid, err = x.Db.GetRepoDidByName(did, name) + if errors.Is(err, sql.ErrNoRows) { + l.Info("repo already torn down or not found", "did", did, "name", name) + w.WriteHeader(http.StatusOK) + return + } + } if err != nil { - fail(xrpcerr.RepoNotFoundError) + l.Error("failed to look up repo", "error", err.Error()) + writeError(w, xrpcerr.GenericError(err), http.StatusInternalServerError) return } - repoPath, _, _, err := x.Db.ResolveRepoDIDOnDisk(x.Config.Repo.ScanPath, repoDid) - if err != nil { - fail(xrpcerr.RepoNotFoundError) + + repoPath, joinErr := securejoin.SecureJoin(x.Config.Repo.ScanPath, repoDid) + if joinErr != nil { + fail(xrpcerr.GenericError(joinErr)) return } @@ -80,22 +93,22 @@ func (x *Xrpc) DeleteRepo(w http.ResponseWriter, r *http.Request) { return } - err = os.RemoveAll(repoPath) - if err != nil { - l.Error("deleting repo", "error", err.Error()) - writeError(w, xrpcerr.GenericError(err), http.StatusInternalServerError) + if rmErr := os.RemoveAll(repoPath); rmErr != nil { + l.Error("deleting repo", "error", rmErr.Error()) + writeError(w, xrpcerr.GenericError(rmErr), http.StatusInternalServerError) return } - err = x.Enforcer.RemoveRepo(did, rbac.ThisServer, repoDid) - if err != nil { - l.Error("failed to delete repo from enforcer", "error", err.Error()) - writeError(w, xrpcerr.GenericError(err), http.StatusInternalServerError) + if rbacErr := x.Enforcer.RemoveRepo(did, rbac.ThisServer, repoDid); rbacErr != nil { + l.Error("failed to delete repo from enforcer", "error", rbacErr.Error()) + writeError(w, xrpcerr.GenericError(rbacErr), http.StatusInternalServerError) return } - if err := x.Db.DeleteRepoKey(repoDid); err != nil { - l.Error("failed to delete repo key", "error", err.Error()) + if delErr := x.Db.DeleteRepoKey(repoDid); delErr != nil { + l.Error("failed to delete repo key", "error", delErr.Error()) + writeError(w, xrpcerr.GenericError(delErr), http.StatusInternalServerError) + return } w.WriteHeader(http.StatusOK)