package xrpc import ( "bytes" "context" "encoding/json" "log/slog" "net/http" "net/http/httptest" "testing" "github.com/bluesky-social/indigo/atproto/syntax" "tangled.org/core/api/org_tangled" "tangled.org/core/rbac" "tangled.org/core/spindle/config" "tangled.org/core/spindle/db" ) func newAllowListXrpc(t *testing.T) (*Xrpc, *db.DB, *rbac.Enforcer, syntax.DID) { t.Helper() d, e := newTestXrpcDB(t) owner := syntax.DID("did:plc:instanceowner") if err := e.AddSpindle(rbac.ThisServer); err != nil { t.Fatalf("AddSpindle: %v", err) } if err := e.AddSpindleOwner(rbac.ThisServer, owner.String()); err != nil { t.Fatalf("AddSpindleOwner: %v", err) } return &Xrpc{ Logger: slog.Default(), Db: d, Enforcer: e, Config: &config.Config{}, }, d, e, owner } func allowListCall(x *Xrpc, handler func(http.ResponseWriter, *http.Request), actor syntax.DID, method, query string, body []byte) *httptest.ResponseRecorder { req := httptest.NewRequest(method, "/"+query, bytes.NewReader(body)) req = req.WithContext(context.WithValue(req.Context(), ActorDid, actor)) w := httptest.NewRecorder() handler(w, req) return w } func TestAllowListRoutesShouldBeOwnerOnly(t *testing.T) { x, _, _, owner := newAllowListXrpc(t) stranger := syntax.DID("did:plc:stranger") for _, call := range []struct { handler func(http.ResponseWriter, *http.Request) method string query string body []byte }{ {x.AllowListAdd, http.MethodPost, org_tangled.TempSpindleAllowListAddNSID, []byte(`{"did":"did:plc:nel"}`)}, {x.AllowListRemove, http.MethodPost, org_tangled.TempSpindleAllowListRemoveNSID, []byte(`{"did":"did:plc:nel"}`)}, {x.AllowListList, http.MethodGet, org_tangled.TempSpindleAllowListListNSID, nil}, } { w := allowListCall(x, call.handler, stranger, call.method, call.query, call.body) if w.Code != http.StatusUnauthorized { t.Fatalf("stranger on %s = %d: %s", call.query, w.Code, w.Body.String()) } } w := allowListCall(x, x.AllowListList, owner, http.MethodGet, org_tangled.TempSpindleAllowListListNSID, nil) if w.Code != http.StatusOK { t.Fatalf("owner list = %d: %s", w.Code, w.Body.String()) } var out org_tangled.TempSpindleAllowListList_Output if err := json.Unmarshal(w.Body.Bytes(), &out); err != nil { t.Fatalf("decode: %v", err) } if len(out.Dids) != 0 { t.Fatalf("A fresh instance lists no entries: %v.", out.Dids) } } func TestAllowListAddAndRemoveShouldRoundTrip(t *testing.T) { x, d, _, owner := newAllowListXrpc(t) nel := syntax.DID("did:plc:nel") squid := syntax.DID("did:plc:squid") for _, did := range []syntax.DID{nel, squid} { body, _ := json.Marshal(map[string]string{"did": did.String()}) w := allowListCall(x, x.AllowListAdd, owner, http.MethodPost, org_tangled.TempSpindleAllowListAddNSID, body) if w.Code != http.StatusOK { t.Fatalf("add %s = %d: %s", did, w.Code, w.Body.String()) } } w := allowListCall(x, x.AllowListList, owner, http.MethodGet, org_tangled.TempSpindleAllowListListNSID, nil) var listed org_tangled.TempSpindleAllowListList_Output if err := json.Unmarshal(w.Body.Bytes(), &listed); err != nil { t.Fatalf("decode: %v", err) } if len(listed.Dids) != 2 || listed.Dids[0] != nel.String() || listed.Dids[1] != squid.String() { t.Fatalf("List = %+v; both entries should return in lexical order.", listed) } body, _ := json.Marshal(map[string]string{"did": squid.String()}) w = allowListCall(x, x.AllowListRemove, owner, http.MethodPost, org_tangled.TempSpindleAllowListRemoveNSID, body) if w.Code != http.StatusOK { t.Fatalf("remove = %d: %s", w.Code, w.Body.String()) } still, err := db.AllowlistContains(d, squid) if err != nil { t.Fatal(err) } if still { t.Fatal("Squid is still on the allow-list after removal.") } } type fakeBackfill struct { tracked []syntax.DID untracked []syntax.DID } func (f *fakeBackfill) Track(dids ...syntax.DID) { f.tracked = append(f.tracked, dids...) } func (f *fakeBackfill) Untrack(did syntax.DID) { f.untracked = append(f.untracked, did) } func TestAllowListChangesShouldReachTheBackfill(t *testing.T) { x, _, _, owner := newAllowListXrpc(t) f := &fakeBackfill{} x.Backfill = f nel := syntax.DID("did:plc:nel") body, _ := json.Marshal(map[string]string{"did": nel.String()}) w := allowListCall(x, x.AllowListAdd, owner, http.MethodPost, org_tangled.TempSpindleAllowListAddNSID, body) if w.Code != http.StatusOK || len(f.tracked) != 1 || f.tracked[0] != nel || len(f.untracked) != 0 { t.Fatalf("add = %d, tracked %v, untracked %v", w.Code, f.tracked, f.untracked) } w = allowListCall(x, x.AllowListRemove, owner, http.MethodPost, org_tangled.TempSpindleAllowListRemoveNSID, body) if w.Code != http.StatusOK || len(f.untracked) != 1 || f.untracked[0] != nel || len(f.tracked) != 1 { t.Fatalf("remove = %d, tracked %v, untracked %v", w.Code, f.tracked, f.untracked) } }