diff --git a/migrations/006_initbans.down.sql b/migrations/006_initbans.down.sql new file mode 100644 index 0000000..d0badcf --- /dev/null +++ b/migrations/006_initbans.down.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS bans; diff --git a/migrations/006_initbans.up.sql b/migrations/006_initbans.up.sql new file mode 100644 index 0000000..e80cb2d --- /dev/null +++ b/migrations/006_initbans.up.sql @@ -0,0 +1,7 @@ +CREATE TABLE bans ( + id SERIAL PRIMARY KEY, + did TEXT PRIMARY KEY, + reason TEXT, + till TIMESTAMPTZ, + banned_at TIMESTAMPTZ NOT NULL DEFAULT now() +); diff --git a/server/internal/db/db.go b/server/internal/db/db.go index 15bd5be..292102e 100644 --- a/server/internal/db/db.go +++ b/server/internal/db/db.go @@ -9,6 +9,7 @@ import ( "rvcx/internal/types" "time" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) @@ -321,8 +322,8 @@ func (s *Store) GetChannelViewHR(handle string, rkey string, ctx context.Context uri := fmt.Sprintf("at://%s/org.xcvr.feed.channel/%s", did, rkey) row := s.pool.QueryRow(ctx, ` SELECT - channels.uri, - channels.host, + channels.uri, + channels.host, channels.title, channels.topic, channels.created_at, @@ -351,3 +352,66 @@ func (s *Store) DeleteChannel(uri string, ctx context.Context) error { _, err := s.pool.Exec(ctx, `DELETE FROM channels WHERE uri = $1`, uri) return err } + +func (s *Store) GetBanned(did string, ctx context.Context) (*types.Ban, error) { + row := s.pool.QueryRow(ctx, `SELECT + id, + reason, + till, + banned_at + FROM bans WHERE did = $1`, did) + var ban types.Ban + err := row.Scan(&ban.Id, &ban.Reason, &ban.Till, &ban.BannedAt) + if err != nil { + return nil, err + } + ban.Did = did + return &ban, nil +} + +func (s *Store) GetBanId(id int, ctx context.Context) (*types.Ban, error) { + row := s.pool.QueryRow(ctx, `SELECT + did, + reason, + till, + banned_at + FROM bans WHERE id = $1`, id) + var ban types.Ban + err := row.Scan(&ban.Id, &ban.Reason, &ban.Till, &ban.BannedAt) + if err != nil { + return nil, err + } + ban.Id = id + return &ban, nil +} + +func (s *Store) AddBan(did string, reason *string, till *time.Time, ctx context.Context) error { + _, err := s.pool.Exec(ctx, `INSERT INTO bans ( + did, + reason, + till + ) VALUES ( + $1, $2, $3 + ) + `, did, reason, till) + return err +} + +func (s *Store) IsBanned(did string, ctx context.Context) (bool, error) { + ban, err := s.GetBanned(did, ctx) + if ban != nil { + defbanned := false + if ban.Till == nil { + defbanned = true + } else { + defbanned = time.Now().Before(*ban.Till) + } + if defbanned { + return true, nil + } + } + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return false, err + } + return false, nil +} diff --git a/server/internal/db/oauth.go b/server/internal/db/oauth.go index 16017f0..f225a39 100644 --- a/server/internal/db/oauth.go +++ b/server/internal/db/oauth.go @@ -101,6 +101,11 @@ func (s Store) DeleteSession(ctx context.Context, did syntax.DID, sessionID stri return nil } +func (s Store) DeleteAllSessions(ctx context.Context, did string) error { + _, err := s.pool.Exec(ctx, `DELETE FROM sessions WHERE account_did = $1`) + return err +} + func (s Store) GetAuthRequestInfo(ctx context.Context, state string) (*oauth.AuthRequestData, error) { row := s.pool.QueryRow(ctx, ` SELECT diff --git a/server/internal/handler/handler.go b/server/internal/handler/handler.go index 36e8fca..b836ddf 100644 --- a/server/internal/handler/handler.go +++ b/server/internal/handler/handler.go @@ -55,6 +55,8 @@ func New(db *db.Store, logger *log.Logger, oauthserv *oauth.Service, model *mode mux.HandleFunc(oauthJWKSPath(), h.WithCORS(h.serveJWKS)) mux.HandleFunc("POST /oauth/login", h.oauthLogin) mux.HandleFunc("POST /oauth/logout", h.oauthMiddleware(h.oauthLogout)) + mux.HandleFunc("POST /oauth/ban", h.postBan) + mux.HandleFunc("GET /oauth/ban", h.getBan) mux.HandleFunc("GET /oauth/whoami", h.getSession) mux.HandleFunc(oauthCallbackPath(), h.WithCORS(h.oauthCallback)) return h diff --git a/server/internal/handler/oauthHandlers.go b/server/internal/handler/oauthHandlers.go index 5260aef..f14fa6f 100644 --- a/server/internal/handler/oauthHandlers.go +++ b/server/internal/handler/oauthHandlers.go @@ -6,8 +6,11 @@ import ( "fmt" "net/http" "os" + "rvcx/internal/atputils" "rvcx/internal/oauth" + "strconv" "strings" + "time" atoauth "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" @@ -57,6 +60,17 @@ func (h *Handler) oauthCallback(w http.ResponseWriter, r *http.Request) { h.serverError(w, errors.New("my god.... :"+err.Error())) return } + isban, err := h.db.IsBanned(sessData.AccountDID.String(), r.Context()) + if err != nil { + h.serverError(w, errors.New("i'm not sure if user is banned, error, "+err.Error())) + return + } + if isban { + ban, _ := h.db.GetBanned(sessData.AccountDID.String(), r.Context()) + http.Redirect(w, r, fmt.Sprintf("%s%d", os.Getenv("BAN_ENDPOINT"), ban.Id), http.StatusSeeOther) + return + } + err = h.rm.CreateInitialProfile(sessData, r.Context()) if err != nil { h.serverError(w, err) @@ -152,3 +166,72 @@ func (h *Handler) oauthMiddleware(f func(cs *atoauth.ClientSession, w http.Respo f(cs, w, r) } } + +func (h *Handler) postBan(w http.ResponseWriter, r *http.Request) { + s, _ := h.sessionStore.Get(r, "oauthsession") + did, bok := s.Values["did"].(string) + if !bok { + h.badRequest(w, errors.New("not authorized")) + return + } + handle, err := h.db.ResolveDid(did, r.Context()) + if err != nil { + h.serverError(w, errors.New("failed to resolve"+err.Error())) + return + } + if handle != os.Getenv("ADMIN_HANDLE") { + h.badRequest(w, errors.New("must be admin to ban")) + return + } + userhandle := r.Header.Get("user") + userdid, err := atputils.GetDidFromHandle(r.Context(), userhandle) + if err != nil { + h.badRequest(w, errors.New("failed to resolve user handle")) + return + } + daysstring := r.Header.Get("days") + daysint, err := strconv.Atoi(daysstring) + var till *time.Time + if err == nil { + tillt := time.Now().Add(time.Hour * 24 * time.Duration(daysint)) + till = &tillt + } + var reason *string + reasonstr := r.Header.Get("reason") + if reasonstr != "" { + reason = &reasonstr + } + err = h.db.AddBan(userdid, reason, till, r.Context()) + if err != nil { + h.serverError(w, errors.New("failed to ban, "+err.Error())) + return + } + ban, err := h.db.GetBanned(userdid, r.Context()) + if err != nil { + h.serverError(w, errors.New("succeeded to ban and then failed again"+err.Error())) + return + } + err = h.db.DeleteAllSessions(r.Context(), ban.Did) + if err != nil { + h.serverError(w, errors.New("failed to kick user "+ban.Did+err.Error())) + return + } + http.Redirect(w, r, fmt.Sprintf("%s%d", os.Getenv("BAN_ENDPOINT"), ban.Id), http.StatusFound) +} + +func (h *Handler) getBan(w http.ResponseWriter, r *http.Request) { + banid := r.Header.Get("id") + id, err := strconv.Atoi(banid) + if err != nil { + h.badRequest(w, err) + return + } + ban, err := h.db.GetBanId(id, r.Context()) + if err != nil { + h.serverError(w, err) + return + } + encoder := json.NewEncoder(w) + w.Header().Add("Content-Type", "application/json") + encoder.Encode(ban) +} diff --git a/server/internal/recordmanager/media.go b/server/internal/recordmanager/media.go index 83cc3fb..850ad25 100644 --- a/server/internal/recordmanager/media.go +++ b/server/internal/recordmanager/media.go @@ -21,8 +21,15 @@ func (rm *RecordManager) PostImage(cs *atoauth.ClientSession, file multipart.Fil } func (rm *RecordManager) AddImageToCache(did string, cid string, ctx context.Context) (string, error) { + ib, err := rm.db.IsBanned(did, ctx) + if err != nil { + return "", err + } + if ib { + return "", errors.New("user banned") + } uploadDir := "./uploads" - _, err := os.Stat(uploadDir) + _, err = os.Stat(uploadDir) if os.IsNotExist(err) { os.Mkdir(uploadDir, 0755) } diff --git a/server/internal/types/oauth.go b/server/internal/types/oauth.go index 15d4c3a..6f44ad9 100644 --- a/server/internal/types/oauth.go +++ b/server/internal/types/oauth.go @@ -29,3 +29,11 @@ type Session struct { RefreshToken string Expiration time.Time } + +type Ban struct { + Id int `json:"id"` + Did string `json:"did"` + Reason *string `json:"reason,omitempty"` + Till *time.Time `json:"till,omitempty"` + BannedAt time.Time `json:"bannedAt"` +}