diff --git a/knotserver/ingester.go b/knotserver/ingester.go --- a/knotserver/ingester.go +++ b/knotserver/ingester.go @@ -73,13 +73,13 @@ } l.Info("added member from firehose", "member", record.Subject) - if err := h.db.AddDid(did); err != nil { + if err := h.db.AddDid(record.Subject); err != nil { l.Error("failed to add did", "error", err) return fmt.Errorf("failed to add did: %w", err) } - h.jc.AddDid(did) + h.jc.AddDid(record.Subject) - if err := h.fetchAndAddKeys(ctx, did); err != nil { + if err := h.fetchAndAddKeys(ctx, record.Subject); err != nil { return fmt.Errorf("failed to fetch and add keys: %w", err) } diff --git a/appview/db/registration.go b/appview/db/registration.go --- a/appview/db/registration.go +++ b/appview/db/registration.go @@ -1,11 +1,8 @@ package db import ( - "crypto/rand" "database/sql" - "encoding/hex" "fmt" - "log" "strings" "time" ) @@ -18,14 +15,29 @@ ByDid string Created *time.Time Registered *time.Time + ReadOnly bool } func (r *Registration) Status() Status { - if r.Registered != nil { + if r.ReadOnly { + return ReadOnly + } else if r.Registered != nil { return Registered } else { return Pending } +} + +func (r *Registration) IsRegistered() bool { + return r.Status() == Registered +} + +func (r *Registration) IsReadOnly() bool { + return r.Status() == ReadOnly +} + +func (r *Registration) IsPending() bool { + return r.Status() == Pending } type Status uint32 @@ -33,161 +45,83 @@ const ( Registered Status = iota Pending + ReadOnly ) -// returns registered status, did of owner, error -func RegistrationsByDid(e Execer, did string) ([]Registration, error) { +func GetRegistrations(e Execer, filters ...filter) ([]Registration, error) { var registrations []Registration - rows, err := e.Query(` - select id, domain, did, created, registered from registrations - where did = ? - `, did) + var conditions []string + var args []any + for _, filter := range filters { + conditions = append(conditions, filter.Condition()) + args = append(args, filter.Arg()...) + } + + whereClause := "" + if conditions != nil { + whereClause = " where " + strings.Join(conditions, " and ") + } + + query := fmt.Sprintf(` + select id, domain, did, created, registered, read_only + from registrations + %s + order by created + `, + whereClause, + ) + + rows, err := e.Query(query, args...) if err != nil { return nil, err } for rows.Next() { - var createdAt *string - var registeredAt *string - var registration Registration - err = rows.Scan(®istration.Id, ®istration.Domain, ®istration.ByDid, &createdAt, ®isteredAt) + var createdAt string + var registeredAt sql.Null[string] + var readOnly int + var reg Registration + err = rows.Scan(®.Id, ®.Domain, ®.ByDid, &createdAt, ®isteredAt, &readOnly) if err != nil { - log.Println(err) - } else { - createdAtTime, _ := time.Parse(time.RFC3339, *createdAt) - var registeredAtTime *time.Time - if registeredAt != nil { - x, _ := time.Parse(time.RFC3339, *registeredAt) - registeredAtTime = &x - } - - registration.Created = &createdAtTime - registration.Registered = registeredAtTime - registrations = append(registrations, registration) + return nil, err } + + if t, err := time.Parse(time.RFC3339, createdAt); err == nil { + reg.Created = &t + } + + if registeredAt.Valid { + if t, err := time.Parse(time.RFC3339, registeredAt.V); err == nil { + reg.Registered = &t + } + } + + if readOnly != 0 { + reg.ReadOnly = true + } + + registrations = append(registrations, reg) } return registrations, nil } -// returns registered status, did of owner, error -func RegistrationByDomain(e Execer, domain string) (*Registration, error) { - var createdAt *string - var registeredAt *string - var registration Registration - - err := e.QueryRow(` - select id, domain, did, created, registered from registrations - where domain = ? - `, domain).Scan(®istration.Id, ®istration.Domain, ®istration.ByDid, &createdAt, ®isteredAt) - - if err != nil { - if err == sql.ErrNoRows { - return nil, nil - } else { - return nil, err - } +func MarkRegistered(e Execer, filters ...filter) error { + var conditions []string + var args []any + for _, filter := range filters { + conditions = append(conditions, filter.Condition()) + args = append(args, filter.Arg()...) } - createdAtTime, _ := time.Parse(time.RFC3339, *createdAt) - var registeredAtTime *time.Time - if registeredAt != nil { - x, _ := time.Parse(time.RFC3339, *registeredAt) - registeredAtTime = &x + query := "update registrations set registered = strftime('%Y-%m-%dT%H:%M:%SZ', 'now'), read_only = 0" + if len(conditions) > 0 { + query += " where " + strings.Join(conditions, " and ") } - registration.Created = &createdAtTime - registration.Registered = registeredAtTime - - return ®istration, nil -} - -func genSecret() string { - key := make([]byte, 32) - rand.Read(key) - return hex.EncodeToString(key) -} - -func GenerateRegistrationKey(e Execer, domain, did string) (string, error) { - // sanity check: does this domain already have a registration? - reg, err := RegistrationByDomain(e, domain) - if err != nil { - return "", err - } - - // registration is open - if reg != nil { - switch reg.Status() { - case Registered: - // already registered by `owner` - return "", fmt.Errorf("%s already registered by %s", domain, reg.ByDid) - case Pending: - // TODO: be loud about this - log.Printf("%s registered by %s, status pending", domain, reg.ByDid) - } - } - - secret := genSecret() - - _, err = e.Exec(` - insert into registrations (domain, did, secret) - values (?, ?, ?) - on conflict(domain) do update set did = excluded.did, secret = excluded.secret, created = excluded.created - `, domain, did, secret) - - if err != nil { - return "", err - } - - return secret, nil -} - -func GetRegistrationKey(e Execer, domain string) (string, error) { - res := e.QueryRow(`select secret from registrations where domain = ?`, domain) - - var secret string - err := res.Scan(&secret) - if err != nil || secret == "" { - return "", err - } - - return secret, nil -} - -func GetCompletedRegistrations(e Execer) ([]string, error) { - rows, err := e.Query(`select domain from registrations where registered not null`) - if err != nil { - return nil, err - } - - var domains []string - for rows.Next() { - var domain string - err = rows.Scan(&domain) - - if err != nil { - log.Println(err) - } else { - domains = append(domains, domain) - } - } - - if err = rows.Err(); err != nil { - return nil, err - } - - return domains, nil -} - -func Register(e Execer, domain string) error { - _, err := e.Exec(` - update registrations - set registered = strftime('%Y-%m-%dT%H:%M:%SZ', 'now') - where domain = ?; - `, domain) - + _, err := e.Exec(query, args...) return err } diff --git a/appview/knots/knots.go b/appview/knots/knots.go --- a/appview/knots/knots.go +++ b/appview/knots/knots.go @@ -3,6 +3,7 @@ import ( "errors" "fmt" + "log" "log/slog" "net/http" "slices" @@ -49,12 +50,17 @@ r.With(middleware.AuthMiddleware(k.OAuth)).Post("/{domain}/add", k.addMember) r.With(middleware.AuthMiddleware(k.OAuth)).Post("/{domain}/remove", k.removeMember) + r.With(middleware.AuthMiddleware(k.OAuth)).Get("/upgradeBanner", k.banner) + return r } func (k *Knots) knots(w http.ResponseWriter, r *http.Request) { user := k.OAuth.GetUser(r) - registrations, err := db.RegistrationsByDid(k.Db, user.Did) + registrations, err := db.GetRegistrations( + k.Db, + db.FilterEq("did", user.Did), + ) if err != nil { k.Logger.Error("failed to fetch knot registrations", "err", err) w.WriteHeader(http.StatusInternalServerError) @@ -89,19 +95,8 @@ http.Error(w, "Not found", http.StatusNotFound) return } - - // Find the specific registration for this domain - var registration *db.Registration - for _, reg := range registrations { - if reg.Domain == domain && reg.ByDid == user.Did && reg.Registered != nil { - registration = ® - break - } - } - - if registration == nil { - l.Error("registration not found or not verified") - http.Error(w, "Not found", http.StatusNotFound) + if len(registrations) != 1 { + l.Error("got incorret number of registrations", "got", len(registrations), "expected", 1) return } registration := registrations[0] @@ -518,10 +513,14 @@ db.FilterIsNot("registered", "null"), ) if err != nil { - l.Error("failed to retrieve domain registration", "err", err) - http.Error(w, "Not found", http.StatusNotFound) + l.Error("failed to get registration", "err", err) return } + if len(registrations) != 1 { + l.Error("got incorret number of registrations", "got", len(registrations), "expected", 1) + return + } + registration := registrations[0] noticeId := fmt.Sprintf("add-member-error-%d", registration.Id) defaultErr := "Failed to add member. Try again later." @@ -678,4 +677,29 @@ // ok k.Pages.HxRefresh(w) +} + +func (k *Knots) banner(w http.ResponseWriter, r *http.Request) { + user := k.OAuth.GetUser(r) + l := k.Logger.With("handler", "removeMember") + l = l.With("did", user.Did) + l = l.With("handle", user.Handle) + + registrations, err := db.GetRegistrations( + k.Db, + db.FilterEq("did", user.Did), + db.FilterEq("read_only", 1), + ) + if err != nil { + l.Error("non-fatal: failed to get registrations") + return + } + + if registrations == nil { + return + } + + k.Pages.KnotBanner(w, pages.KnotBannerParams{ + Registrations: registrations, + }) } diff --git a/appview/pages/pages.go b/appview/pages/pages.go --- a/appview/pages/pages.go +++ b/appview/pages/pages.go @@ -338,6 +338,14 @@ return p.execute("user/settings/emails", w, params) } +type KnotBannerParams struct { + Registrations []db.Registration +} + +func (p *Pages) KnotBanner(w io.Writer, params KnotBannerParams) error { + return p.executePlain("knots/fragments/banner", w, params) +} + type KnotsParams struct { LoggedInUser *oauth.User Registrations []db.Registration @@ -360,7 +368,7 @@ } type KnotListingParams struct { - db.Registration + *db.Registration } func (p *Pages) KnotListing(w io.Writer, params KnotListingParams) error { diff --git a/appview/state/knotstream.go b/appview/state/knotstream.go --- a/appview/state/knotstream.go +++ b/appview/state/knotstream.go @@ -24,14 +24,17 @@ ) func Knotstream(ctx context.Context, c *config.Config, d *db.DB, enforcer *rbac.Enforcer, posthog posthog.Client) (*ec.Consumer, error) { - knots, err := db.GetCompletedRegistrations(d) + knots, err := db.GetRegistrations( + d, + db.FilterIsNot("registered", "null"), + ) if err != nil { return nil, err } srcs := make(map[ec.Source]struct{}) for _, k := range knots { - s := ec.NewKnotSource(k) + s := ec.NewKnotSource(k.Domain) srcs[s] = struct{}{} } diff --git a/appview/state/state.go b/appview/state/state.go --- a/appview/state/state.go +++ b/appview/state/state.go @@ -435,9 +435,9 @@ Rkey: rkey, }, ) - if err != nil { - l.Error("xrpc request failed", "err", err) - s.pages.Notice(w, "repo", fmt.Sprintf("Failed to create repository on knot server: %s.", err.Error())) + if err := xrpcclient.HandleXrpcErr(xe); err != nil { + l.Error("xrpc error", "xe", xe) + s.pages.Notice(w, "repo", err.Error()) return } diff --git a/appview/xrpcclient/xrpc.go b/appview/xrpcclient/xrpc.go --- a/appview/xrpcclient/xrpc.go +++ b/appview/xrpcclient/xrpc.go @@ -3,10 +3,14 @@ import ( "bytes" "context" + "errors" + "fmt" "io" + "net/http" "github.com/bluesky-social/indigo/api/atproto" "github.com/bluesky-social/indigo/xrpc" + indigoxrpc "github.com/bluesky-social/indigo/xrpc" oauth "tangled.sh/icyphox.sh/atproto-oauth" ) @@ -101,4 +105,25 @@ } return &out, nil +} + +// produces a more manageable error +func HandleXrpcErr(err error) error { + if err == nil { + return nil + } + + var xrpcerr *indigoxrpc.Error + if ok := errors.As(err, &xrpcerr); !ok { + return fmt.Errorf("Recieved invalid XRPC error response.") + } + + switch xrpcerr.StatusCode { + case http.StatusNotFound: + return fmt.Errorf("XRPC is unsupported on this knot, consider upgrading your knot.") + case http.StatusUnauthorized: + return fmt.Errorf("Unauthorized XRPC request.") + default: + return fmt.Errorf("Failed to perform operation. Try again later.") + } } diff --git a/appview/pages/templates/knots/index.html b/appview/pages/templates/knots/index.html --- a/appview/pages/templates/knots/index.html +++ b/appview/pages/templates/knots/index.html @@ -77,7 +77,7 @@ -
+ diff --git a/appview/pages/templates/layouts/topbar.html b/appview/pages/templates/layouts/topbar.html --- a/appview/pages/templates/layouts/topbar.html +++ b/appview/pages/templates/layouts/topbar.html @@ -21,6 +21,13 @@ + {{ if .LoggedInUser }} + + {{ end }} {{ end }} {{ define "newButton" }} diff --git a/appview/pages/templates/knots/fragments/banner.html b/appview/pages/templates/knots/fragments/banner.html new file mode 100644 --- /dev/null +++ b/appview/pages/templates/knots/fragments/banner.html @@ -0,0 +1,9 @@ +{{ define "knots/fragments/banner" }} +