466 lines
11 KiB
Go
466 lines
11 KiB
Go
package store
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
// Memory is an in-process Store for tests.
|
|
type Memory struct {
|
|
mu sync.Mutex
|
|
users map[string]*User // id -> user
|
|
byName map[string]string // username -> id
|
|
posts map[string]*Post // id -> post
|
|
postVotes map[string]map[string]int // postID -> userID -> value
|
|
}
|
|
|
|
// NewMemory returns an empty Memory store.
|
|
func NewMemory() *Memory {
|
|
return &Memory{
|
|
users: map[string]*User{},
|
|
byName: map[string]string{},
|
|
posts: map[string]*Post{},
|
|
postVotes: map[string]map[string]int{},
|
|
}
|
|
}
|
|
|
|
func (m *Memory) CreateUser(_ context.Context, u *User) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if u.Role != RoleUser && u.Role != RoleAdmin {
|
|
return fmt.Errorf("invalid role")
|
|
}
|
|
u.Username = NormalizeUsername(u.Username)
|
|
u.Email = NormalizeEmail(u.Email)
|
|
if _, ok := m.byName[u.Username]; ok {
|
|
return ErrDuplicateUsername
|
|
}
|
|
if u.Email != "" {
|
|
for _, existing := range m.users {
|
|
if existing.Email == u.Email {
|
|
return ErrDuplicateEmail
|
|
}
|
|
}
|
|
}
|
|
if u.ID == "" {
|
|
u.ID = uuid.NewString()
|
|
}
|
|
if u.Name == "" {
|
|
u.Name = u.Username
|
|
}
|
|
if u.CreatedAt == "" {
|
|
u.CreatedAt = time.Now().UTC().Format(time.RFC3339)
|
|
}
|
|
role := u.Role
|
|
if role == RoleAdmin {
|
|
for _, existing := range m.users {
|
|
if existing.Role == RoleAdmin {
|
|
role = RoleUser
|
|
break
|
|
}
|
|
}
|
|
}
|
|
cp := *u
|
|
cp.Role = role
|
|
cp.db = nil
|
|
m.users[cp.ID] = &cp
|
|
m.byName[cp.Username] = cp.ID
|
|
*u = cp
|
|
return nil
|
|
}
|
|
|
|
func (m *Memory) UserByID(_ context.Context, id string) (*User, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
u, ok := m.users[id]
|
|
if !ok {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
cp := *u
|
|
return &cp, nil
|
|
}
|
|
|
|
func (m *Memory) UserByUsername(_ context.Context, username string) (*User, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
id, ok := m.byName[NormalizeUsername(username)]
|
|
if !ok {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
cp := *m.users[id]
|
|
return &cp, nil
|
|
}
|
|
|
|
func (m *Memory) ListUsers(_ context.Context, q ListUsersQuery) ([]User, string, string, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
limit := q.Limit
|
|
if limit <= 0 {
|
|
limit = AdminUsersLimit
|
|
}
|
|
search := strings.ToLower(strings.TrimSpace(q.Search))
|
|
out := make([]User, 0, len(m.users))
|
|
for _, u := range m.users {
|
|
if search != "" &&
|
|
!strings.Contains(strings.ToLower(u.Username), search) &&
|
|
!strings.Contains(strings.ToLower(u.Name), search) {
|
|
continue
|
|
}
|
|
if q.CursorCreated != "" {
|
|
if u.CreatedAt > q.CursorCreated {
|
|
continue
|
|
}
|
|
if u.CreatedAt == q.CursorCreated && u.ID >= q.CursorID {
|
|
continue
|
|
}
|
|
}
|
|
out = append(out, *u)
|
|
}
|
|
sort.Slice(out, func(i, j int) bool {
|
|
if out[i].CreatedAt != out[j].CreatedAt {
|
|
return out[i].CreatedAt > out[j].CreatedAt
|
|
}
|
|
return out[i].ID > out[j].ID
|
|
})
|
|
var nextCreated, nextID string
|
|
if len(out) > limit {
|
|
last := out[limit-1]
|
|
nextCreated, nextID = last.CreatedAt, last.ID
|
|
out = out[:limit]
|
|
}
|
|
return out, nextCreated, nextID, nil
|
|
}
|
|
|
|
func (m *Memory) CountAdmins(_ context.Context) (int, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
n := 0
|
|
for _, u := range m.users {
|
|
if u.Role == RoleAdmin {
|
|
n++
|
|
}
|
|
}
|
|
return n, nil
|
|
}
|
|
|
|
// SetUserRole serializes demotions under m.mu (same critical section as count).
|
|
func (m *Memory) SetUserRole(_ context.Context, id string, role Role) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if role != RoleUser && role != RoleAdmin {
|
|
return fmt.Errorf("invalid role")
|
|
}
|
|
u, ok := m.users[id]
|
|
if !ok {
|
|
return sql.ErrNoRows
|
|
}
|
|
if u.Role == RoleAdmin && role == RoleUser {
|
|
n := 0
|
|
for _, x := range m.users {
|
|
if x.Role == RoleAdmin {
|
|
n++
|
|
}
|
|
}
|
|
if n <= 1 {
|
|
return ErrLastAdmin
|
|
}
|
|
}
|
|
u.Role = role
|
|
return nil
|
|
}
|
|
|
|
func (m *Memory) SaveUserProfile(_ context.Context, u *User) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
cur, ok := m.users[u.ID]
|
|
if !ok {
|
|
return sql.ErrNoRows
|
|
}
|
|
email := NormalizeEmail(u.Email)
|
|
if email != "" {
|
|
for id, existing := range m.users {
|
|
if id != u.ID && existing.Email == email {
|
|
return ErrDuplicateEmail
|
|
}
|
|
}
|
|
}
|
|
cur.State = strings.TrimSpace(u.State)
|
|
cur.Email = email
|
|
if u.AvatarURL != "" {
|
|
cur.AvatarURL = u.AvatarURL
|
|
}
|
|
u.State = cur.State
|
|
u.Email = cur.Email
|
|
u.AvatarURL = cur.AvatarURL
|
|
return nil
|
|
}
|
|
|
|
func (m *Memory) CreatePost(_ context.Context, post *Post) error {
|
|
if post == nil {
|
|
return fmt.Errorf("%w: post is nil", ErrInvalidPost)
|
|
}
|
|
if err := preparePost(post); err != nil {
|
|
return err
|
|
}
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if _, ok := m.users[post.AuthorID]; !ok {
|
|
return fmt.Errorf("%w: unknown author", ErrInvalidPost)
|
|
}
|
|
if _, exists := m.posts[post.ID]; exists {
|
|
return fmt.Errorf("%w: duplicate id", ErrInvalidPost)
|
|
}
|
|
if post.ParentID != nil {
|
|
if _, ok := m.posts[*post.ParentID]; !ok {
|
|
return fmt.Errorf("%w: unknown parent", ErrInvalidPost)
|
|
}
|
|
}
|
|
cp := clonePost(post)
|
|
cp.db = nil
|
|
m.posts[cp.ID] = cp
|
|
*post = *clonePost(cp)
|
|
return nil
|
|
}
|
|
|
|
func (m *Memory) GetPost(_ context.Context, id string) (*Post, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
post, ok := m.posts[id]
|
|
if !ok {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
return clonePostWithAuthor(post, m.users), nil
|
|
}
|
|
|
|
func (m *Memory) GetPostThread(_ context.Context, rootID string) (*Post, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
root, ok := m.posts[rootID]
|
|
if !ok || root.ParentID != nil {
|
|
return nil, sql.ErrNoRows
|
|
}
|
|
inThread := map[string]bool{rootID: true}
|
|
for changed := true; changed; {
|
|
changed = false
|
|
for id, post := range m.posts {
|
|
if inThread[id] || post.ParentID == nil || !inThread[*post.ParentID] {
|
|
continue
|
|
}
|
|
inThread[id] = true
|
|
changed = true
|
|
}
|
|
}
|
|
posts := make([]Post, 0, len(inThread))
|
|
for id := range inThread {
|
|
posts = append(posts, *clonePostWithAuthor(m.posts[id], m.users))
|
|
}
|
|
return buildPostTree(posts, rootID)
|
|
}
|
|
|
|
func (m *Memory) GetPostThreadForViewer(ctx context.Context, rootID, viewerID string) (*Post, error) {
|
|
root, err := m.GetPostThread(ctx, rootID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
for _, value := range m.postVotes[rootID] {
|
|
root.Score += value
|
|
}
|
|
root.UserVote = m.postVotes[rootID][viewerID]
|
|
root.Answered = m.threadContainsAdminReply(rootID)
|
|
return root, nil
|
|
}
|
|
|
|
func (m *Memory) UpdatePost(_ context.Context, post *Post) error {
|
|
if post == nil {
|
|
return fmt.Errorf("%w: post is nil", ErrInvalidPost)
|
|
}
|
|
body := strings.TrimSpace(post.Body)
|
|
if body == "" {
|
|
return fmt.Errorf("%w: body is required", ErrInvalidPost)
|
|
}
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
existing, ok := m.posts[post.ID]
|
|
if !ok {
|
|
return sql.ErrNoRows
|
|
}
|
|
existing.Body = body
|
|
existing.UpdatedAt = time.Now().UTC().Format(time.RFC3339Nano)
|
|
*post = *clonePostWithAuthor(existing, m.users)
|
|
return nil
|
|
}
|
|
|
|
func (m *Memory) ListRootPosts(_ context.Context, postDate, viewerID string) ([]Post, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
posts := make([]Post, 0)
|
|
for _, post := range m.posts {
|
|
if post.ParentID != nil || post.PostDate != postDate || post.PostState == PostStateHidden {
|
|
continue
|
|
}
|
|
cp := clonePostWithAuthor(post, m.users)
|
|
for _, value := range m.postVotes[post.ID] {
|
|
cp.Score += value
|
|
}
|
|
cp.UserVote = m.postVotes[post.ID][viewerID]
|
|
cp.Answered = m.threadContainsAdminReply(post.ID)
|
|
posts = append(posts, *cp)
|
|
}
|
|
sort.Slice(posts, func(i, j int) bool {
|
|
if posts[i].Score != posts[j].Score {
|
|
return posts[i].Score > posts[j].Score
|
|
}
|
|
if posts[i].CreatedAt != posts[j].CreatedAt {
|
|
return posts[i].CreatedAt < posts[j].CreatedAt
|
|
}
|
|
return posts[i].ID < posts[j].ID
|
|
})
|
|
if len(posts) > HuntListLimit {
|
|
posts = posts[:HuntListLimit]
|
|
}
|
|
return posts, nil
|
|
}
|
|
|
|
func (m *Memory) ListRootPostsByAuthor(_ context.Context, authorID string) ([]Post, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
posts := make([]Post, 0)
|
|
for _, post := range m.posts {
|
|
if post.ParentID != nil ||
|
|
post.AuthorID != authorID ||
|
|
post.PostState == PostStateHidden {
|
|
continue
|
|
}
|
|
posts = append(posts, *clonePostWithAuthor(post, m.users))
|
|
}
|
|
return sortProfilePosts(posts), nil
|
|
}
|
|
|
|
func (m *Memory) ListRootPostsAnsweredBy(_ context.Context, adminID string) ([]Post, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
posts := make([]Post, 0)
|
|
for _, root := range m.posts {
|
|
if root.ParentID != nil || root.PostState == PostStateHidden {
|
|
continue
|
|
}
|
|
participated := false
|
|
for _, post := range m.posts {
|
|
if post.AuthorID == adminID && m.postIsDescendantOf(post, root.ID) {
|
|
participated = true
|
|
break
|
|
}
|
|
}
|
|
if participated {
|
|
posts = append(posts, *clonePostWithAuthor(root, m.users))
|
|
}
|
|
}
|
|
return sortProfilePosts(posts), nil
|
|
}
|
|
|
|
func (m *Memory) SetRootPostState(_ context.Context, id string, state PostState) error {
|
|
switch state {
|
|
case PostStateVisible, PostStateHidden, PostStateLocked:
|
|
default:
|
|
return fmt.Errorf("%w: invalid post state", ErrInvalidPost)
|
|
}
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
post, ok := m.posts[id]
|
|
if !ok || post.ParentID != nil {
|
|
return sql.ErrNoRows
|
|
}
|
|
post.PostState = state
|
|
post.UpdatedAt = time.Now().UTC().Format(time.RFC3339Nano)
|
|
return nil
|
|
}
|
|
|
|
func sortProfilePosts(posts []Post) []Post {
|
|
sort.Slice(posts, func(i, j int) bool {
|
|
if posts[i].CreatedAt != posts[j].CreatedAt {
|
|
return posts[i].CreatedAt > posts[j].CreatedAt
|
|
}
|
|
return posts[i].ID > posts[j].ID
|
|
})
|
|
if len(posts) > ProfileListLimit {
|
|
posts = posts[:ProfileListLimit]
|
|
}
|
|
return posts
|
|
}
|
|
|
|
func (m *Memory) VotePost(_ context.Context, userID, postID string, value int) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if value != 1 && value != -1 && value != 0 {
|
|
return fmt.Errorf("invalid vote")
|
|
}
|
|
post, ok := m.posts[postID]
|
|
if !ok || post.ParentID != nil || post.PostState == PostStateHidden {
|
|
return ErrPostNotVotable
|
|
}
|
|
if m.postVotes[postID] == nil {
|
|
m.postVotes[postID] = map[string]int{}
|
|
}
|
|
if value == 0 {
|
|
delete(m.postVotes[postID], userID)
|
|
return nil
|
|
}
|
|
m.postVotes[postID][userID] = value
|
|
return nil
|
|
}
|
|
|
|
func (m *Memory) threadContainsAdminReply(rootID string) bool {
|
|
for id, post := range m.posts {
|
|
if id == rootID || !m.postIsDescendantOf(post, rootID) {
|
|
continue
|
|
}
|
|
if author := m.users[post.AuthorID]; author != nil && author.Role == RoleAdmin {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func (m *Memory) postIsDescendantOf(post *Post, rootID string) bool {
|
|
seen := map[string]bool{}
|
|
for post != nil && post.ParentID != nil {
|
|
if *post.ParentID == rootID {
|
|
return true
|
|
}
|
|
if seen[*post.ParentID] {
|
|
return false
|
|
}
|
|
seen[*post.ParentID] = true
|
|
post = m.posts[*post.ParentID]
|
|
}
|
|
return false
|
|
}
|
|
|
|
func clonePost(post *Post) *Post {
|
|
cp := *post
|
|
if post.ParentID != nil {
|
|
parentID := *post.ParentID
|
|
cp.ParentID = &parentID
|
|
}
|
|
cp.Replies = nil
|
|
return &cp
|
|
}
|
|
|
|
func clonePostWithAuthor(post *Post, users map[string]*User) *Post {
|
|
cp := clonePost(post)
|
|
if author := users[post.AuthorID]; author != nil {
|
|
cp.AuthorName = author.Name
|
|
cp.AuthorRole = author.Role
|
|
}
|
|
return cp
|
|
}
|