diff --git a/Server/api/block_handler.go b/Server/api/block_handler.go deleted file mode 100644 index 9ac661d0..00000000 --- a/Server/api/block_handler.go +++ /dev/null @@ -1,137 +0,0 @@ -package api - -import ( - "log/slog" - "net/http" - "strconv" - - "github.com/go-chi/chi/v5" - "github.com/owncord/server/db" -) - -// handleBlockUser blocks a user (prevents DM creation and messaging). -func handleBlockUser(database *db.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - user, ok := r.Context().Value(UserKey).(*db.User) - if !ok || user == nil { - writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "not authenticated", - }) - return - } - - targetID, err := strconv.ParseInt(chi.URLParam(r, "userId"), 10, 64) - if err != nil || targetID <= 0 { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "invalid user ID", - }) - return - } - - if targetID == user.ID { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "cannot block yourself", - }) - return - } - - // Verify target user exists. - target, err := database.GetUserByID(targetID) - if err != nil { - slog.Error("handleBlockUser GetUserByID", "err", err, "target_id", targetID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up user", - }) - return - } - if target == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "user not found", - }) - return - } - - if err := database.BlockUser(user.ID, targetID); err != nil { - slog.Error("handleBlockUser BlockUser", "err", err, - "blocker_id", user.ID, "blocked_id", targetID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to block user", - }) - return - } - - slog.Info("user blocked", "blocker_id", user.ID, "blocked_id", targetID) - writeJSON(w, http.StatusOK, map[string]string{"message": "user blocked"}) - } -} - -// handleUnblockUser removes a block on a user. -func handleUnblockUser(database *db.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - user, ok := r.Context().Value(UserKey).(*db.User) - if !ok || user == nil { - writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "not authenticated", - }) - return - } - - targetID, err := strconv.ParseInt(chi.URLParam(r, "userId"), 10, 64) - if err != nil || targetID <= 0 { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "invalid user ID", - }) - return - } - - if err := database.UnblockUser(user.ID, targetID); err != nil { - slog.Error("handleUnblockUser UnblockUser", "err", err, - "blocker_id", user.ID, "target_id", targetID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to unblock user", - }) - return - } - - slog.Info("user unblocked", "blocker_id", user.ID, "unblocked_id", targetID) - writeJSON(w, http.StatusOK, map[string]string{"message": "user unblocked"}) - } -} - -// handleListBlocks returns all users blocked by the authenticated user. -func handleListBlocks(database *db.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - user, ok := r.Context().Value(UserKey).(*db.User) - if !ok || user == nil { - writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "not authenticated", - }) - return - } - - ids, err := database.ListBlockedUsers(user.ID) - if err != nil { - slog.Error("handleListBlocks ListBlockedUsers", "err", err, "user_id", user.ID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to list blocked users", - }) - return - } - if ids == nil { - ids = []int64{} - } - - writeJSON(w, http.StatusOK, map[string]any{"blocked_user_ids": ids}) - } -} diff --git a/Server/api/channel_handler.go b/Server/api/channel_handler.go index 03d6b6c1..5b556d7a 100644 --- a/Server/api/channel_handler.go +++ b/Server/api/channel_handler.go @@ -1,6 +1,7 @@ package api import ( + "errors" "log/slog" "net/http" "strconv" @@ -10,7 +11,7 @@ import ( "github.com/go-chi/chi/v5" "github.com/owncord/server/auth" "github.com/owncord/server/db" - "github.com/owncord/server/permissions" + "github.com/owncord/server/service" ) const ( @@ -41,7 +42,6 @@ func searchRateLimitMiddleware(limiter *auth.RateLimiter, limit int, window time }) return } - next.ServeHTTP(w, r) }) } @@ -50,163 +50,66 @@ func searchRateLimitMiddleware(limiter *auth.RateLimiter, limit int, window time // MountChannelRoutes registers all channel-related routes onto r. // All routes require authentication. The limiter is used to rate-limit // expensive endpoints like search. -func MountChannelRoutes(r chi.Router, database *db.DB, limiter *auth.RateLimiter, trustedProxies []string) { +func MountChannelRoutes(r chi.Router, database *db.DB, svc *service.Services, limiter *auth.RateLimiter, trustedProxies []string) { r.Route("/api/v1/channels", func(r chi.Router) { r.Use(AuthMiddleware(database)) - r.Get("/", handleListChannels(database)) - r.Get("/{id}/messages", handleGetMessages(database)) - r.Get("/{id}/pins", handleGetPins(database)) - r.Post("/{id}/pins/{messageId}", handleSetPinned(database, true)) - r.Delete("/{id}/pins/{messageId}", handleSetPinned(database, false)) + r.Get("/", handleListChannels(svc)) + r.Get("/{id}/messages", handleGetMessages(svc)) + r.Get("/{id}/pins", handleGetPins(svc)) + r.Post("/{id}/pins/{messageId}", handleSetPinned(svc, true)) + r.Delete("/{id}/pins/{messageId}", handleSetPinned(svc, false)) }) r.With( AuthMiddleware(database), searchRateLimitMiddleware(limiter, searchRateLimitPerMinute, time.Minute, trustedProxies), - ).Get("/api/v1/search", handleSearch(database)) -} - -// hasChannelPermREST checks whether the role has the given permission on the channel, -// accounting for Administrator bypass and channel overrides. -func hasChannelPermREST(database *db.DB, role *db.Role, channelID, perm int64) bool { - if role == nil { - return false - } - if permissions.HasAdmin(role.Permissions) { - return true - } - allow, deny, err := database.GetChannelPermissions(channelID, role.ID) - if err != nil { - return false - } - effective := permissions.EffectivePerms(role.Permissions, allow, deny) - return effective&perm == perm -} - -// hasChannelPermBatch checks permission using a pre-fetched overrides map, -// eliminating N+1 queries when filtering multiple channels. -func hasChannelPermBatch(role *db.Role, overrides map[int64]db.ChannelOverride, channelID, perm int64) bool { - if role == nil { - return false - } - if permissions.HasAdmin(role.Permissions) { - return true - } - o := overrides[channelID] // zero-value (0,0) when no override exists - effective := permissions.EffectivePerms(role.Permissions, o.Allow, o.Deny) - return effective&perm == perm + ).Get("/api/v1/search", handleSearch(svc)) } // handleListChannels returns all channels the authenticated user can see. -func handleListChannels(database *db.DB) http.HandlerFunc { +func handleListChannels(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - role, _ := r.Context().Value(RoleKey).(*db.Role) - - channels, err := database.ListChannels() - if err != nil { - slog.Error("handleListChannels ListChannels", "err", err) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to list channels", + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", Message: "authentication required", }) return } - // Batch-fetch all channel permission overrides for this role in one query. - overrides := map[int64]db.ChannelOverride{} - if role != nil && !permissions.HasAdmin(role.Permissions) { - var oErr error - overrides, oErr = database.GetAllChannelPermissionsForRole(role.ID) - if oErr != nil { - slog.Error("handleListChannels GetAllChannelPermissionsForRole", "err", oErr) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to fetch channel permissions", - }) - return - } + channels, err := svc.Channels.ListVisibleChannels(user.ID) + if err != nil { + slog.Error("handleListChannels", "err", err) + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "INTERNAL_ERROR", Message: "failed to list channels", + }) + return } - - // Filter channels by READ_MESSAGES permission. - // DM channels are excluded — they are delivered via the separate DM endpoints. - var visible []db.Channel - for i := range channels { - if channels[i].Type == "dm" { - continue - } - if hasChannelPermBatch(role, overrides, channels[i].ID, permissions.ReadMessages) { - visible = append(visible, channels[i]) - } - } - if visible == nil { - visible = []db.Channel{} - } - writeJSON(w, http.StatusOK, visible) + writeJSON(w, http.StatusOK, channels) } } // handleGetMessages returns paginated messages for a channel. -// Query params: before (int64, message ID for pagination), limit (1-100, default 50). -func handleGetMessages(database *db.DB) http.HandlerFunc { +func handleGetMessages(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { channelID, ok := parseIDParam(w, r, "id") if !ok { return } - ch, err := database.GetChannel(channelID) - if err != nil { - slog.Error("handleGetMessages GetChannel", "err", err, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up channel", - }) - return - } - if ch == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", Message: "authentication required", }) return } - // DM channels use participant-based auth instead of role-based permissions. - if ch.Type == "dm" { - user, _ := r.Context().Value(UserKey).(*db.User) - if user == nil { - writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "authentication required", - }) - return - } - ok, dmErr := database.IsDMParticipant(user.ID, channelID) - if dmErr != nil || !ok { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", - }) - return - } - } else { - role, _ := r.Context().Value(RoleKey).(*db.Role) - if !hasChannelPermREST(database, role, channelID, permissions.ReadMessages) { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "FORBIDDEN", - Message: "no permission to view this channel", - }) - return - } - } - - // Parse query params. before := int64(0) if raw := r.URL.Query().Get("before"); raw != "" { v, parseErr := strconv.ParseInt(raw, 10, 64) if parseErr != nil || v < 0 { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "before must be a non-negative integer", + Error: "BAD_REQUEST", Message: "before must be a non-negative integer", }) return } @@ -218,8 +121,7 @@ func handleGetMessages(database *db.DB) http.HandlerFunc { v, parseErr := strconv.Atoi(raw) if parseErr != nil || v < 1 { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "limit must be a positive integer", + Error: "BAD_REQUEST", Message: "limit must be a positive integer", }) return } @@ -229,29 +131,12 @@ func handleGetMessages(database *db.DB) http.HandlerFunc { limit = v } - // Extract requesting user ID for reaction "me" flag. - var userID int64 - if user, ok := r.Context().Value(UserKey).(*db.User); ok && user != nil { - userID = user.ID - } - - // Fetch one extra to determine has_more. - msgs, err := database.GetMessagesForAPI(channelID, before, limit+1, userID) + msgs, hasMore, err := svc.Messages.GetMessages(user.ID, channelID, before, limit) if err != nil { - slog.Error("handleGetMessages GetMessagesForAPI", "err", err, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to fetch messages", - }) + writeServiceError(w, err) return } - hasMore := false - if len(msgs) > limit { - hasMore = true - msgs = msgs[:limit] - } - type response struct { Messages []db.MessageAPIResponse `json:"messages"` HasMore bool `json:"has_more"` @@ -261,14 +146,20 @@ func handleGetMessages(database *db.DB) http.HandlerFunc { } // handleSearch performs a full-text search across messages. -// Query params: q (required), channel_id (optional), limit (optional, 1-100). -func handleSearch(database *db.DB) http.HandlerFunc { +func handleSearch(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { q := r.URL.Query().Get("q") if q == "" { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "query parameter 'q' is required", + Error: "BAD_REQUEST", Message: "query parameter 'q' is required", + }) + return + } + + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", Message: "authentication required", }) return } @@ -278,47 +169,11 @@ func handleSearch(database *db.DB) http.HandlerFunc { v, parseErr := strconv.ParseInt(raw, 10, 64) if parseErr != nil || v <= 0 { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "channel_id must be a positive integer", + Error: "BAD_REQUEST", Message: "channel_id must be a positive integer", }) return } channelID = &v - - // Pre-check: verify the user can read this channel before running - // the FTS query, preventing timing-oracle information leakage. - ch, chErr := database.GetChannel(v) - if chErr != nil || ch == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", - }) - return - } - if ch.Type == "dm" { - user, _ := r.Context().Value(UserKey).(*db.User) - if user == nil { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "FORBIDDEN", Message: "no permission to search this channel", - }) - return - } - ok, dmErr := database.IsDMParticipant(user.ID, v) - if dmErr != nil || !ok { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "FORBIDDEN", Message: "no permission to search this channel", - }) - return - } - } else { - role, _ := r.Context().Value(RoleKey).(*db.Role) - if !hasChannelPermREST(database, role, v, permissions.ReadMessages) { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "FORBIDDEN", Message: "no permission to search this channel", - }) - return - } - } } limit := defaultMessageLimit @@ -326,8 +181,7 @@ func handleSearch(database *db.DB) http.HandlerFunc { v, parseErr := strconv.Atoi(raw) if parseErr != nil || v < 1 { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "limit must be a positive integer", + Error: "BAD_REQUEST", Message: "limit must be a positive integer", }) return } @@ -337,105 +191,19 @@ func handleSearch(database *db.DB) http.HandlerFunc { limit = v } - var results []db.MessageSearchResult - - if channelID != nil { - // Single-channel search: permission already checked above. - var err error - results, err = database.SearchMessages(q, channelID, limit) - if err != nil { - if isInvalidSearchQueryError(err) { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "invalid search query", - }) - return - } - slog.Error("handleSearch SearchMessages", "err", err, "query", q) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "search failed", + results, err := svc.Messages.SearchMessages(user.ID, q, channelID, limit) + if err != nil { + if isInvalidSearchQueryError(err) { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "BAD_REQUEST", Message: "invalid search query", }) return } - } else { - // Global search: pre-compute the set of accessible channel IDs - // so the DB query never touches restricted content. - role, _ := r.Context().Value(RoleKey).(*db.Role) - user, _ := r.Context().Value(UserKey).(*db.User) - - // 1. Guild channels the user can read. - allChannels, chErr := database.ListChannels() - if chErr != nil { - slog.Error("handleSearch ListChannels", "err", chErr) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "search failed", - }) - return - } - - overrides := map[int64]db.ChannelOverride{} - if role != nil && !permissions.HasAdmin(role.Permissions) { - var oErr error - overrides, oErr = database.GetAllChannelPermissionsForRole(role.ID) - if oErr != nil { - slog.Error("handleSearch GetAllChannelPermissionsForRole", "err", oErr) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "search failed", - }) - return - } - } - - var accessibleIDs []int64 - for i := range allChannels { - if allChannels[i].Type == "dm" { - continue // DM channels handled separately below. - } - if hasChannelPermBatch(role, overrides, allChannels[i].ID, permissions.ReadMessages) { - accessibleIDs = append(accessibleIDs, allChannels[i].ID) - } - } - - // 2. DM channels the user participates in. - if user != nil { - dmChannels, dmErr := database.GetUserDMChannels(user.ID) - if dmErr != nil { - slog.Error("handleSearch GetUserDMChannels", "err", dmErr) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "search failed", - }) - return - } - for _, dm := range dmChannels { - accessibleIDs = append(accessibleIDs, dm.ChannelID) - } - } - - if len(accessibleIDs) == 0 { - results = []db.MessageSearchResult{} - } else { - var err error - results, err = database.SearchMessagesInChannels(q, accessibleIDs, limit) - if err != nil { - if isInvalidSearchQueryError(err) { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "invalid search query", - }) - return - } - slog.Error("handleSearch SearchMessagesInChannels", "err", err, "query", q) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "search failed", - }) - return - } - } + writeServiceError(w, err) + return + } + if results == nil { + results = []db.MessageSearchResult{} } type response struct { @@ -446,73 +214,24 @@ func handleSearch(database *db.DB) http.HandlerFunc { } // handleGetPins returns all pinned messages for a channel. -func handleGetPins(database *db.DB) http.HandlerFunc { +func handleGetPins(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { channelID, ok := parseIDParam(w, r, "id") if !ok { return } - ch, err := database.GetChannel(channelID) - if err != nil { - slog.Error("handleGetPins GetChannel", "err", err, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up channel", - }) - return - } - if ch == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", Message: "authentication required", }) return } - // DM channels use participant-based auth instead of role-based permissions. - if ch.Type == "dm" { - user, _ := r.Context().Value(UserKey).(*db.User) - if user == nil { - writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "authentication required", - }) - return - } - ok, dmErr := database.IsDMParticipant(user.ID, channelID) - if dmErr != nil || !ok { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", - }) - return - } - } else { - // Permission check: user must have READ_MESSAGES on this channel. - role, _ := r.Context().Value(RoleKey).(*db.Role) - if !hasChannelPermREST(database, role, channelID, permissions.ReadMessages) { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "FORBIDDEN", - Message: "no permission to view this channel", - }) - return - } - } - - // Extract requesting user ID for reaction "me" flag. - var userID int64 - if user, ok := r.Context().Value(UserKey).(*db.User); ok && user != nil { - userID = user.ID - } - - msgs, err := database.GetPinnedMessages(channelID, userID) + msgs, err := svc.Messages.GetPinnedMessages(user.ID, channelID) if err != nil { - slog.Error("handleGetPins GetPinnedMessages", "err", err, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to fetch pinned messages", - }) + writeServiceError(w, err) return } @@ -525,101 +244,52 @@ func handleGetPins(database *db.DB) http.HandlerFunc { } // handleSetPinned pins or unpins a message in a channel. -func handleSetPinned(database *db.DB, pinned bool) http.HandlerFunc { - action := "pin" - if !pinned { - action = "unpin" - } +func handleSetPinned(svc *service.Services, pinned bool) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { channelID, ok := parseIDParam(w, r, "id") if !ok { return } - messageID, ok := parseIDParam(w, r, "messageId") if !ok { return } - // Look up the channel to check if it's a DM. - ch, chErr := database.GetChannel(channelID) - if chErr != nil { - slog.Error("handleSetPinned GetChannel", "err", chErr, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up channel", - }) - return - } - if ch == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", Message: "authentication required", }) return } - // DM channels use participant-based auth instead of role-based permissions. - if ch.Type == "dm" { - user, _ := r.Context().Value(UserKey).(*db.User) - if user == nil { - writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "authentication required", - }) - return - } - ok, dmErr := database.IsDMParticipant(user.ID, channelID) - if dmErr != nil || !ok { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", - }) - return - } - } else { - // Permission check: user must have MANAGE_MESSAGES on this channel. - role, _ := r.Context().Value(RoleKey).(*db.Role) - if !hasChannelPermREST(database, role, channelID, permissions.ManageMessages) { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "FORBIDDEN", - Message: "no permission to manage messages in this channel", - }) - return - } - } - - // Verify message exists and belongs to this channel. - msg, err := database.GetMessage(messageID) - if err != nil { - slog.Error("handleSetPinned GetMessage", "err", err, "action", action, "message_id", messageID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up message", - }) + if err := svc.Messages.SetMessagePinned(user.ID, channelID, messageID, pinned); err != nil { + writeServiceError(w, err) return } - if msg == nil || msg.ChannelID != channelID { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "message not found", - }) - return - } - - if err := database.SetMessagePinned(messageID, pinned); err != nil { - slog.Error("handleSetPinned SetMessagePinned", "err", err, "action", action, "message_id", messageID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to " + action + " message", - }) - return - } - w.WriteHeader(http.StatusNoContent) } } +// writeServiceError maps a service-layer error to an HTTP response. +func writeServiceError(w http.ResponseWriter, err error) { + switch { + case errors.Is(err, service.ErrRateLimited): + writeJSON(w, http.StatusTooManyRequests, errorResponse{Error: "RATE_LIMITED", Message: err.Error()}) + case errors.Is(err, service.ErrBadRequest): + writeJSON(w, http.StatusBadRequest, errorResponse{Error: "BAD_REQUEST", Message: err.Error()}) + case errors.Is(err, service.ErrNotFound): + writeJSON(w, http.StatusNotFound, errorResponse{Error: "NOT_FOUND", Message: err.Error()}) + case errors.Is(err, service.ErrForbidden), errors.Is(err, service.ErrBlocked): + writeJSON(w, http.StatusForbidden, errorResponse{Error: "FORBIDDEN", Message: err.Error()}) + case errors.Is(err, service.ErrConflict): + writeJSON(w, http.StatusConflict, errorResponse{Error: "CONFLICT", Message: err.Error()}) + default: + slog.Error("service error", "err", err) + writeJSON(w, http.StatusInternalServerError, errorResponse{Error: "INTERNAL_ERROR", Message: "internal error"}) + } +} + // parseIDParam extracts and validates a chi URL param as int64. // Writes a 400 response and returns false on failure. func parseIDParam(w http.ResponseWriter, r *http.Request, param string) (int64, bool) { diff --git a/Server/api/dm_handler.go b/Server/api/dm_handler.go index fbe79989..fef4e44e 100644 --- a/Server/api/dm_handler.go +++ b/Server/api/dm_handler.go @@ -8,6 +8,7 @@ import ( "github.com/go-chi/chi/v5" "github.com/owncord/server/db" + "github.com/owncord/server/service" ) // DMBroadcaster is the interface needed to send WebSocket events from REST @@ -19,20 +20,20 @@ type DMBroadcaster interface { // MountDMRoutes registers DM-related routes onto r. // All routes require authentication. // hub is used to send real-time WebSocket events on DM close. -func MountDMRoutes(r chi.Router, database *db.DB, broadcaster DMBroadcaster) { +func MountDMRoutes(r chi.Router, database *db.DB, svc *service.Services, broadcaster DMBroadcaster) { r.Route("/api/v1/dms", func(r chi.Router) { r.Use(AuthMiddleware(database)) - r.Post("/", handleCreateDM(database)) - r.Get("/", handleListDMs(database)) - r.Delete("/{channelId}", handleCloseDM(database, broadcaster)) + r.Post("/", handleCreateDM(svc)) + r.Get("/", handleListDMs(svc)) + r.Delete("/{channelId}", handleCloseDM(svc, broadcaster)) }) // User blocking routes — prevent DM creation and messaging. r.Route("/api/v1/blocks", func(r chi.Router) { r.Use(AuthMiddleware(database)) - r.Get("/", handleListBlocks(database)) - r.Put("/{userId}", handleBlockUser(database)) - r.Delete("/{userId}", handleUnblockUser(database)) + r.Get("/", handleListBlocks(svc)) + r.Put("/{userId}", handleBlockUser(svc)) + r.Delete("/{userId}", handleUnblockUser(svc)) }) } @@ -54,13 +55,12 @@ type listDMsResponse struct { } // handleCreateDM creates or retrieves a DM channel with a recipient. -func handleCreateDM(database *db.DB) http.HandlerFunc { +func handleCreateDM(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { user, ok := r.Context().Value(UserKey).(*db.User) if !ok || user == nil { writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "authentication required", + Error: "UNAUTHORIZED", Message: "authentication required", }) return } @@ -68,137 +68,67 @@ func handleCreateDM(database *db.DB) http.HandlerFunc { var req createDMRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "invalid request body", + Error: "BAD_REQUEST", Message: "invalid request body", }) return } - if req.RecipientID <= 0 { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "recipient_id must be a positive integer", - }) - return - } - - // Cannot DM yourself. - if req.RecipientID == user.ID { - writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "cannot create a DM with yourself", - }) - return - } - - // Verify recipient exists. - recipient, err := database.GetUserByID(req.RecipientID) + result, err := svc.DMs.CreateDM(user.ID, req.RecipientID) if err != nil { - slog.Error("handleCreateDM GetUserByID", "err", err, "recipient_id", req.RecipientID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up recipient", - }) - return - } - if recipient == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "recipient not found", - }) + writeServiceError(w, err) return } - // Check if either user has blocked the other. - blocked, blockErr := database.IsEitherBlocked(user.ID, req.RecipientID) - if blockErr != nil { - slog.Error("handleCreateDM IsEitherBlocked", "err", blockErr, - "user_id", user.ID, "recipient_id", req.RecipientID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to check block status", - }) - return - } - if blocked { - writeJSON(w, http.StatusForbidden, errorResponse{ - Error: "BLOCKED", - Message: "cannot create DM with this user", - }) - return - } - - // Get or create the DM channel. - ch, created, err := database.GetOrCreateDMChannel(user.ID, req.RecipientID) //nolint:contextcheck // TODO: propagate context through this call path - if err != nil { - slog.Error("handleCreateDM GetOrCreateDMChannel", "err", err, - "user_id", user.ID, "recipient_id", req.RecipientID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to create DM channel", - }) - return - } - - // Build the recipient DMUser from the fetched user. avatarStr := "" - if recipient.Avatar != nil { - avatarStr = *recipient.Avatar + if result.Recipient.Avatar != nil { + avatarStr = *result.Recipient.Avatar } dmUser := db.DMUser{ - ID: recipient.ID, - Username: recipient.Username, + ID: result.Recipient.ID, + Username: result.Recipient.Username, Avatar: avatarStr, - Status: recipient.Status, + Status: result.Recipient.Status, } status := http.StatusOK - if created { + if result.Created { status = http.StatusCreated } - writeJSON(w, status, createDMResponse{ - ChannelID: ch.ID, + ChannelID: result.Channel.ID, Recipient: dmUser, - Created: created, + Created: result.Created, }) } } // handleListDMs returns all open DM channels for the authenticated user. -func handleListDMs(database *db.DB) http.HandlerFunc { +func handleListDMs(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { user, ok := r.Context().Value(UserKey).(*db.User) if !ok || user == nil { writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "authentication required", + Error: "UNAUTHORIZED", Message: "authentication required", }) return } - channels, err := database.GetUserDMChannels(user.ID) + channels, err := svc.DMs.ListDMs(user.ID) if err != nil { - slog.Error("handleListDMs GetUserDMChannels", "err", err, "user_id", user.ID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to list DM channels", - }) + writeServiceError(w, err) return } - writeJSON(w, http.StatusOK, listDMsResponse{DMChannels: channels}) } } // handleCloseDM removes a DM channel from the authenticated user's open list. -func handleCloseDM(database *db.DB, broadcaster DMBroadcaster) http.HandlerFunc { +func handleCloseDM(svc *service.Services, broadcaster DMBroadcaster) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { user, ok := r.Context().Value(UserKey).(*db.User) if !ok || user == nil { writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "authentication required", + Error: "UNAUTHORIZED", Message: "authentication required", }) return } @@ -208,45 +138,81 @@ func handleCloseDM(database *db.DB, broadcaster DMBroadcaster) http.HandlerFunc return } - // Verify user is a participant in this DM. - isParticipant, err := database.IsDMParticipant(user.ID, channelID) - if err != nil { - slog.Error("handleCloseDM IsDMParticipant", "err", err, - "user_id", user.ID, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to verify DM participation", - }) - return - } - if !isParticipant { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "channel not found", - }) + if err := svc.DMs.CloseDM(user.ID, channelID); err != nil { + writeServiceError(w, err) return } - if err := database.CloseDM(user.ID, channelID); err != nil { - slog.Error("handleCloseDM CloseDM", "err", err, - "user_id", user.ID, "channel_id", channelID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to close DM", - }) - return - } - - // Notify the closing user's WebSocket connections so the sidebar updates - // immediately without waiting for a reconnect. + // Notify via WebSocket so sidebar updates immediately. if broadcaster != nil { closeMsg := []byte(fmt.Sprintf(`{"type":"dm_channel_close","payload":{"channel_id":%d}}`, channelID)) if ok := broadcaster.SendToUser(user.ID, closeMsg); !ok { - slog.Debug("handleCloseDM: user not connected, WS notify skipped", - "user_id", user.ID, "channel_id", channelID) + slog.Debug("handleCloseDM: user not connected", "user_id", user.ID, "channel_id", channelID) } } w.WriteHeader(http.StatusNoContent) } } + +// handleBlockUser blocks a user. +func handleBlockUser(svc *service.Services) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{Error: "UNAUTHORIZED", Message: "authentication required"}) + return + } + + targetID, ok := parseIDParam(w, r, "userId") + if !ok { + return + } + + if err := svc.Blocks.BlockUser(user.ID, targetID); err != nil { + writeServiceError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]string{"message": "user blocked"}) + } +} + +// handleUnblockUser unblocks a user. +func handleUnblockUser(svc *service.Services) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{Error: "UNAUTHORIZED", Message: "authentication required"}) + return + } + + targetID, ok := parseIDParam(w, r, "userId") + if !ok { + return + } + + if err := svc.Blocks.UnblockUser(user.ID, targetID); err != nil { + writeServiceError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]string{"message": "user unblocked"}) + } +} + +// handleListBlocks returns all blocked user IDs. +func handleListBlocks(svc *service.Services) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + user, _ := r.Context().Value(UserKey).(*db.User) + if user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{Error: "UNAUTHORIZED", Message: "authentication required"}) + return + } + + ids, err := svc.Blocks.ListBlocked(user.ID) + if err != nil { + writeServiceError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]any{"blocked_user_ids": ids}) + } +} diff --git a/Server/api/invite_handler.go b/Server/api/invite_handler.go index dbb7b077..4030e810 100644 --- a/Server/api/invite_handler.go +++ b/Server/api/invite_handler.go @@ -4,13 +4,12 @@ import ( "encoding/json" "fmt" "io" - "log/slog" "net/http" - "time" "github.com/go-chi/chi/v5" "github.com/owncord/server/db" "github.com/owncord/server/permissions" + "github.com/owncord/server/service" ) // createInviteRequest is the JSON body for POST /api/v1/invites. @@ -32,28 +31,25 @@ type inviteResponse struct { // MountInviteRoutes registers invite endpoints on the given router. // All routes require authentication and MANAGE_INVITES permission. -func MountInviteRoutes(r chi.Router, database *db.DB) { +func MountInviteRoutes(r chi.Router, database *db.DB, svc *service.Services) { r.Route("/api/v1/invites", func(r chi.Router) { r.Use(AuthMiddleware(database)) r.Use(RequirePermission(permissions.ManageInvites)) - r.Post("/", handleCreateInvite(database)) - r.Get("/", handleListInvites(database)) - r.Delete("/{code}", handleRevokeInvite(database)) + r.Post("/", handleCreateInvite(svc)) + r.Get("/", handleListInvites(svc)) + r.Delete("/{code}", handleRevokeInvite(svc)) }) } // handleCreateInvite processes POST /api/v1/invites. -func handleCreateInvite(database *db.DB) http.HandlerFunc { +func handleCreateInvite(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var req createInviteRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - // An empty body is valid (all fields optional), but malformed - // JSON must be rejected so callers notice typos. if err != io.EOF { writeJSON(w, http.StatusBadRequest, errorResponse{ - Error: "BAD_REQUEST", - Message: "malformed JSON body", + Error: "BAD_REQUEST", Message: "malformed JSON body", }) return } @@ -63,62 +59,35 @@ func handleCreateInvite(database *db.DB) http.HandlerFunc { user, ok := r.Context().Value(UserKey).(*db.User) if !ok || user == nil { writeJSON(w, http.StatusUnauthorized, errorResponse{ - Error: "UNAUTHORIZED", - Message: "not authenticated", + Error: "UNAUTHORIZED", Message: "not authenticated", }) return } - var expiresAt *time.Time - // H-4: Cap invite expiration to 30 days (720 hours) to prevent - // effectively permanent invites that survive admin revocation policies. - const maxExpiresInHours = 720 - if req.ExpiresInHours > maxExpiresInHours { + // H-4: Cap invite expiration to 30 days. + if req.ExpiresInHours > service.MaxInviteExpiryHours() { writeJSON(w, http.StatusBadRequest, errorResponse{ Error: "BAD_REQUEST", - Message: fmt.Sprintf("expires_in_hours cannot exceed %d (30 days)", maxExpiresInHours), + Message: fmt.Sprintf("expires_in_hours cannot exceed %d (30 days)", service.MaxInviteExpiryHours()), }) return } - if req.ExpiresInHours > 0 { - t := time.Now().Add(time.Duration(req.ExpiresInHours) * time.Hour) - expiresAt = &t - } - code, err := database.CreateInvite(user.ID, req.MaxUses, expiresAt) + inv, err := svc.Invites.CreateInvite(user.ID, req.MaxUses, req.ExpiresInHours) if err != nil { - slog.Error("handleCreateInvite CreateInvite", "err", err, "user_id", user.ID) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to create invite", - }) + writeServiceError(w, err) return } - - inv, err := database.GetInvite(code) - if err != nil || inv == nil { - slog.Error("handleCreateInvite GetInvite", "err", err, "code", code) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to retrieve invite", - }) - return - } - writeJSON(w, http.StatusCreated, toInviteResponse(inv)) } } // handleListInvites processes GET /api/v1/invites. -func handleListInvites(database *db.DB) http.HandlerFunc { +func handleListInvites(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - invites, err := database.ListInvites() + invites, err := svc.Invites.ListInvites() if err != nil { - slog.Error("handleListInvites ListInvites", "err", err) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to list invites", - }) + writeServiceError(w, err) return } @@ -131,36 +100,13 @@ func handleListInvites(database *db.DB) http.HandlerFunc { } // handleRevokeInvite processes DELETE /api/v1/invites/:code. -func handleRevokeInvite(database *db.DB) http.HandlerFunc { +func handleRevokeInvite(svc *service.Services) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { code := chi.URLParam(r, "code") - - inv, err := database.GetInvite(code) - if err != nil { - slog.Error("handleRevokeInvite GetInvite", "err", err, "code", code) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to look up invite", - }) + if err := svc.Invites.RevokeInvite(code); err != nil { + writeServiceError(w, err) return } - if inv == nil { - writeJSON(w, http.StatusNotFound, errorResponse{ - Error: "NOT_FOUND", - Message: "invite not found", - }) - return - } - - if err := database.RevokeInvite(code); err != nil { - slog.Error("handleRevokeInvite RevokeInvite", "err", err, "code", code) - writeJSON(w, http.StatusInternalServerError, errorResponse{ - Error: "INTERNAL_ERROR", - Message: "failed to revoke invite", - }) - return - } - w.WriteHeader(http.StatusNoContent) } } diff --git a/Server/api/router.go b/Server/api/router.go index f8cbcd0b..9d81183f 100644 --- a/Server/api/router.go +++ b/Server/api/router.go @@ -18,6 +18,7 @@ import ( "github.com/owncord/server/permissions" "github.com/owncord/server/service" "github.com/owncord/server/storage" + dbstore "github.com/owncord/server/store" "github.com/owncord/server/updater" "github.com/owncord/server/ws" ) @@ -90,10 +91,10 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri // broadcast user_update events for real-time profile changes. // Invite management routes (require MANAGE_INVITES permission). - MountInviteRoutes(r, database) + MountInviteRoutes(r, database, svc) // Channel and message REST routes. - MountChannelRoutes(r, database, limiter, cfg.Server.TrustedProxies) + MountChannelRoutes(r, database, svc, limiter, cfg.Server.TrustedProxies) // DM REST routes are mounted after hub creation (below) so the hub can // be passed as a DMBroadcaster for real-time close events. @@ -113,7 +114,8 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri } // Service layer — centralizes business logic for REST and WS handlers. - svc := service.New(database, limiter) + st := dbstore.NewSQLiteStore(database) + svc := service.New(st, limiter) // WebSocket hub — WS does its own in-band auth, so no AuthMiddleware here. hub := ws.NewHub(database, limiter, svc) @@ -179,7 +181,7 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri // DM (direct message) REST routes — mounted after hub creation so the // hub can send real-time dm_channel_close events to WebSocket clients. - MountDMRoutes(r, database, hub) + MountDMRoutes(r, database, svc, hub) // H-8: Connectivity diagnostics restricted to admin users only. // Exposes Go runtime version and LiveKit node IP which aid targeted attacks. diff --git a/Server/service/block.go b/Server/service/block.go new file mode 100644 index 00000000..898250c9 --- /dev/null +++ b/Server/service/block.go @@ -0,0 +1,65 @@ +package service + +import ( + "fmt" + "log/slog" + + "github.com/owncord/server/store" +) + +// BlockService handles user block/unblock operations. +type BlockService struct { + st store.Store +} + +// NewBlockService creates a BlockService. +func NewBlockService(st store.Store) *BlockService { + return &BlockService{st: st} +} + +// BlockUser blocks a target user. Validates the target exists and +// prevents self-blocking. +func (s *BlockService) BlockUser(blockerID, targetID int64) error { + if targetID <= 0 { + return fmt.Errorf("%w: user_id must be positive", ErrBadRequest) + } + if blockerID == targetID { + return fmt.Errorf("%w: cannot block yourself", ErrBadRequest) + } + + target, err := s.st.GetUserByID(targetID) + if err != nil || target == nil { + return fmt.Errorf("%w: user not found", ErrNotFound) + } + + if err := s.st.BlockUser(blockerID, targetID); err != nil { + return fmt.Errorf("%w: failed to block user", ErrInternal) + } + + slog.Info("user blocked", "blocker_id", blockerID, "target_id", targetID) + return nil +} + +// UnblockUser removes a block on a target user. +func (s *BlockService) UnblockUser(blockerID, targetID int64) error { + if targetID <= 0 { + return fmt.Errorf("%w: user_id must be positive", ErrBadRequest) + } + if err := s.st.UnblockUser(blockerID, targetID); err != nil { + return fmt.Errorf("%w: failed to unblock user", ErrInternal) + } + slog.Info("user unblocked", "blocker_id", blockerID, "target_id", targetID) + return nil +} + +// ListBlocked returns all user IDs blocked by the given user. +func (s *BlockService) ListBlocked(blockerID int64) ([]int64, error) { + ids, err := s.st.ListBlockedUsers(blockerID) + if err != nil { + return nil, fmt.Errorf("%w: failed to list blocked users", ErrInternal) + } + if ids == nil { + ids = []int64{} + } + return ids, nil +} diff --git a/Server/service/channel.go b/Server/service/channel.go index 097e1e89..c6e84f93 100644 --- a/Server/service/channel.go +++ b/Server/service/channel.go @@ -7,19 +7,20 @@ import ( "github.com/owncord/server/db" "github.com/owncord/server/permissions" + "github.com/owncord/server/store" ) // ChannelService handles channel-related business logic including // listing, permission-filtered access, typing, presence, and read state. type ChannelService struct { - db *db.DB + st store.Store perms *PermissionService } // NewChannelService creates a ChannelService. -func NewChannelService(database *db.DB, perms *PermissionService) *ChannelService { +func NewChannelService(st store.Store, perms *PermissionService) *ChannelService { return &ChannelService{ - db: database, + st: st, perms: perms, } } @@ -27,7 +28,7 @@ func NewChannelService(database *db.DB, perms *PermissionService) *ChannelServic // ListVisibleChannels returns channels the user has ReadMessages permission for. // DM channels are excluded (they are accessed via DMService). func (s *ChannelService) ListVisibleChannels(userID int64) ([]db.Channel, error) { - all, err := s.db.ListChannels() + all, err := s.st.ListChannels() if err != nil { slog.Error("ChannelService.ListVisibleChannels", "err", err) return nil, fmt.Errorf("%w: failed to list channels", ErrInternal) @@ -51,7 +52,7 @@ func (s *ChannelService) ListVisibleChannels(userID int64) ([]db.Channel, error) return visible, nil } - overrides, err := s.db.GetAllChannelPermissionsForRole(role.ID) + overrides, err := s.st.GetAllChannelPermissionsForRole(role.ID) if err != nil { overrides = make(map[int64]db.ChannelOverride) } @@ -90,13 +91,13 @@ func (s *ChannelService) HandleTyping(userID, channelID int64, limiter interface return nil, nil } - ch, err := s.db.GetChannel(channelID) + ch, err := s.st.GetChannel(channelID) if err != nil || ch == nil { return nil, nil // silent drop } if ch.Type == "dm" { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil || !ok { return nil, nil // silent drop } @@ -110,7 +111,7 @@ func (s *ChannelService) HandleTyping(userID, channelID int64, limiter interface // GetDMParticipantIDs returns the participant IDs for a DM channel. // Convenience method for handlers building DM events. func (s *ChannelService) GetDMParticipantIDs(channelID int64) ([]int64, error) { - return s.db.GetDMParticipantIDs(channelID) + return s.st.GetDMParticipantIDs(channelID) } // HandlePresenceUpdate validates and persists a presence status change. @@ -130,7 +131,7 @@ func (s *ChannelService) HandlePresenceUpdate(userID int64, status string, limit return fmt.Errorf("%w: invalid status", ErrBadRequest) } - if err := s.db.UpdateUserStatus(userID, status); err != nil { + if err := s.st.UpdateUserStatus(userID, status); err != nil { slog.Error("ChannelService.HandlePresenceUpdate", "err", err, "user_id", userID) return fmt.Errorf("%w: failed to update status", ErrInternal) } @@ -145,13 +146,13 @@ func (s *ChannelService) HandleChannelFocus(userID, channelID int64) (*db.Channe return nil, fmt.Errorf("%w: channel_id must be positive", ErrBadRequest) } - ch, err := s.db.GetChannel(channelID) + ch, err := s.st.GetChannel(channelID) if err != nil || ch == nil { return nil, fmt.Errorf("%w: channel not found", ErrForbidden) } if ch.Type == "dm" { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil || !ok { return nil, fmt.Errorf("%w: access denied", ErrForbidden) } @@ -162,9 +163,9 @@ func (s *ChannelService) HandleChannelFocus(userID, channelID int64) (*db.Channe } // Mark channel as read. - latestID, err := s.db.GetLatestMessageID(channelID) + latestID, err := s.st.GetLatestMessageID(channelID) if err == nil && latestID > 0 { - _ = s.db.UpdateReadState(userID, channelID, latestID) + _ = s.st.UpdateReadState(userID, channelID, latestID) } slog.Debug("channel_focus", "user_id", userID, "channel_id", channelID) diff --git a/Server/service/dm.go b/Server/service/dm.go new file mode 100644 index 00000000..e6cd32c3 --- /dev/null +++ b/Server/service/dm.go @@ -0,0 +1,90 @@ +package service + +import ( + "fmt" + "log/slog" + + "github.com/owncord/server/db" + "github.com/owncord/server/store" +) + +// DMService handles direct message channel operations. +type DMService struct { + st store.Store +} + +// NewDMService creates a DMService. +func NewDMService(st store.Store) *DMService { + return &DMService{st: st} +} + +// CreateDMResult holds the result of creating or fetching a DM channel. +type CreateDMResult struct { + Channel *db.Channel + Created bool + Recipient *db.User +} + +// CreateDM creates or retrieves a DM channel between two users. +// Validates that neither user has blocked the other. +func (s *DMService) CreateDM(userID, recipientID int64) (*CreateDMResult, error) { + if recipientID <= 0 { + return nil, fmt.Errorf("%w: recipient_id must be positive", ErrBadRequest) + } + if userID == recipientID { + return nil, fmt.Errorf("%w: cannot create DM with yourself", ErrBadRequest) + } + + recipient, err := s.st.GetUserByID(recipientID) + if err != nil || recipient == nil { + return nil, fmt.Errorf("%w: recipient not found", ErrNotFound) + } + + blocked, err := s.st.IsEitherBlocked(userID, recipientID) + if err != nil { + return nil, fmt.Errorf("%w: failed to check block status", ErrInternal) + } + if blocked { + return nil, fmt.Errorf("%w: cannot create DM — user is blocked", ErrForbidden) + } + + ch, created, err := s.st.GetOrCreateDMChannel(userID, recipientID) + if err != nil { + slog.Error("DMService.CreateDM", "err", err) + return nil, fmt.Errorf("%w: failed to create DM channel", ErrInternal) + } + + return &CreateDMResult{ + Channel: ch, + Created: created, + Recipient: recipient, + }, nil +} + +// ListDMs returns all open DM channels for a user. +func (s *DMService) ListDMs(userID int64) ([]db.DMChannelInfo, error) { + dms, err := s.st.GetUserDMChannels(userID) + if err != nil { + return nil, fmt.Errorf("%w: failed to list DMs", ErrInternal) + } + return dms, nil +} + +// CloseDM closes a DM channel for a user. +func (s *DMService) CloseDM(userID, channelID int64) error { + if channelID <= 0 { + return fmt.Errorf("%w: channel_id must be positive", ErrBadRequest) + } + + ok, err := s.st.IsDMParticipant(userID, channelID) + if err != nil || !ok { + return fmt.Errorf("%w: not a participant in this DM", ErrForbidden) + } + + if err := s.st.CloseDM(userID, channelID); err != nil { + return fmt.Errorf("%w: failed to close DM", ErrInternal) + } + + slog.Debug("DM closed", "user_id", userID, "channel_id", channelID) + return nil +} diff --git a/Server/service/invite.go b/Server/service/invite.go new file mode 100644 index 00000000..a4b79547 --- /dev/null +++ b/Server/service/invite.go @@ -0,0 +1,71 @@ +package service + +import ( + "fmt" + "time" + + "github.com/owncord/server/db" + "github.com/owncord/server/store" +) + +// InviteService handles invite management. +type InviteService struct { + st store.Store +} + +// NewInviteService creates an InviteService. +func NewInviteService(st store.Store) *InviteService { + return &InviteService{st: st} +} + +// maxInviteExpiryHoursVal caps invite expiry to 30 days (H-4 hardening). +const maxInviteExpiryHoursVal = 720 + +// MaxInviteExpiryHours returns the maximum invite expiry in hours. +func MaxInviteExpiryHours() int { return maxInviteExpiryHoursVal } + +// CreateInvite creates a new invite code with optional max uses and expiry. +func (s *InviteService) CreateInvite(createdBy int64, maxUses int, expiresInHours int) (*db.Invite, error) { + // Cap expiry. + if expiresInHours > maxInviteExpiryHoursVal { + expiresInHours = maxInviteExpiryHoursVal + } + + var expiresAt *time.Time + if expiresInHours > 0 { + t := time.Now().Add(time.Duration(expiresInHours) * time.Hour) + expiresAt = &t + } + + code, err := s.st.CreateInvite(createdBy, maxUses, expiresAt) + if err != nil { + return nil, fmt.Errorf("%w: failed to create invite", ErrInternal) + } + + invite, err := s.st.GetInvite(code) + if err != nil { + return nil, fmt.Errorf("%w: failed to fetch invite", ErrInternal) + } + return invite, nil +} + +// ListInvites returns all invites. +func (s *InviteService) ListInvites() ([]*db.Invite, error) { + invites, err := s.st.ListInvites() + if err != nil { + return nil, fmt.Errorf("%w: failed to list invites", ErrInternal) + } + return invites, nil +} + +// RevokeInvite revokes an invite by code. +func (s *InviteService) RevokeInvite(code string) error { + invite, err := s.st.GetInvite(code) + if err != nil || invite == nil { + return fmt.Errorf("%w: invite not found", ErrNotFound) + } + if err := s.st.RevokeInvite(code); err != nil { + return fmt.Errorf("%w: failed to revoke invite", ErrInternal) + } + return nil +} diff --git a/Server/service/message.go b/Server/service/message.go index d1bc4489..0a78ee15 100644 --- a/Server/service/message.go +++ b/Server/service/message.go @@ -11,6 +11,7 @@ import ( "github.com/owncord/server/auth" "github.com/owncord/server/db" "github.com/owncord/server/permissions" + "github.com/owncord/server/store" ) // sanitizer is the shared HTML sanitization policy (strips all tags). @@ -97,15 +98,15 @@ type ReactionResult struct { // MessageService handles message-related business logic including // send, edit, delete, reactions, pins, and search. type MessageService struct { - db *db.DB + st store.Store perms *PermissionService limiter *auth.RateLimiter } // NewMessageService creates a MessageService. -func NewMessageService(database *db.DB, perms *PermissionService, limiter *auth.RateLimiter) *MessageService { +func NewMessageService(st store.Store, perms *PermissionService, limiter *auth.RateLimiter) *MessageService { return &MessageService{ - db: database, + st: st, perms: perms, limiter: limiter, } @@ -124,7 +125,7 @@ func (s *MessageService) SendMessage(p SendMessageParams) (*SendMessageResult, e return nil, fmt.Errorf("%w: channel_id must be a positive integer", ErrBadRequest) } - ch, err := s.db.GetChannel(p.ChannelID) + ch, err := s.st.GetChannel(p.ChannelID) if err != nil || ch == nil { return nil, fmt.Errorf("%w: channel not found", ErrNotFound) } @@ -158,7 +159,7 @@ func (s *MessageService) SendMessage(p SendMessageParams) (*SendMessageResult, e } // Persist message. - msgID, err := s.db.CreateMessage(p.ChannelID, p.UserID, content, p.ReplyTo) + msgID, err := s.st.CreateMessage(p.ChannelID, p.UserID, content, p.ReplyTo) if err != nil { slog.Error("MessageService.SendMessage CreateMessage", "err", err) return nil, fmt.Errorf("%w: failed to save message", ErrInternal) @@ -167,17 +168,17 @@ func (s *MessageService) SendMessage(p SendMessageParams) (*SendMessageResult, e // Link attachments. var attachments []db.AttachmentInfo if len(p.AttachmentIDs) > 0 { - linked, linkErr := s.db.LinkAttachmentsToMessage(msgID, p.AttachmentIDs) + linked, linkErr := s.st.LinkAttachmentsToMessage(msgID, p.AttachmentIDs) if linkErr != nil { slog.Error("MessageService.SendMessage LinkAttachments", "err", linkErr, "msg_id", msgID) // Cleanup: soft-delete the message. - if delErr := s.db.DeleteMessage(msgID, p.UserID, true); delErr != nil { + if delErr := s.st.DeleteMessage(msgID, p.UserID, true); delErr != nil { slog.Error("MessageService.SendMessage DeleteMessage (cleanup)", "err", delErr, "msg_id", msgID) } return nil, fmt.Errorf("%w: failed to send message with attachments", ErrInternal) } if linked > 0 { - attMap, attErr := s.db.GetAttachmentsByMessageIDs([]int64{msgID}) + attMap, attErr := s.st.GetAttachmentsByMessageIDs([]int64{msgID}) if attErr != nil { slog.Error("MessageService.SendMessage GetAttachments", "err", attErr) } else { @@ -187,7 +188,7 @@ func (s *MessageService) SendMessage(p SendMessageParams) (*SendMessageResult, e } // Fetch message for timestamp. - msg, err := s.db.GetMessage(msgID) + msg, err := s.st.GetMessage(msgID) if err != nil || msg == nil { slog.Error("MessageService.SendMessage GetMessage after create", "err", err) return nil, fmt.Errorf("%w: failed to retrieve message", ErrInternal) @@ -204,21 +205,21 @@ func (s *MessageService) SendMessage(p SendMessageParams) (*SendMessageResult, e // DM path: open DM for recipients. if isDM { - participantIDs, pErr := s.db.GetDMParticipantIDs(p.ChannelID) + participantIDs, pErr := s.st.GetDMParticipantIDs(p.ChannelID) if pErr != nil { slog.Error("MessageService.SendMessage GetDMParticipantIDs", "err", pErr, "channel_id", p.ChannelID) return result, nil // Message saved, skip DM side effects. } result.ParticipantIDs = participantIDs - sender, _ := s.db.GetUserByID(p.UserID) + sender, _ := s.st.GetUserByID(p.UserID) result.SenderUser = sender for _, pid := range participantIDs { if pid == p.UserID { continue } - if openErr := s.db.OpenDM(pid, p.ChannelID); openErr != nil { + if openErr := s.st.OpenDM(pid, p.ChannelID); openErr != nil { slog.Error("MessageService.SendMessage OpenDM", "err", openErr, "recipient_id", pid, "channel_id", p.ChannelID) continue } @@ -248,7 +249,7 @@ func (s *MessageService) EditMessage(userID, msgID int64, rawContent string) (*E } // Fetch message. - msg, err := s.db.GetMessage(msgID) + msg, err := s.st.GetMessage(msgID) if err != nil || msg == nil { return nil, fmt.Errorf("%w: cannot edit this message", ErrForbidden) } @@ -257,11 +258,11 @@ func (s *MessageService) EditMessage(userID, msgID int64, rawContent string) (*E } // Channel type for DM-aware permissions. - ch, chErr := s.db.GetChannel(msg.ChannelID) + ch, chErr := s.st.GetChannel(msg.ChannelID) isDM := chErr == nil && ch != nil && ch.Type == "dm" if isDM { - ok, dmErr := s.db.IsDMParticipant(userID, msg.ChannelID) + ok, dmErr := s.st.IsDMParticipant(userID, msg.ChannelID) if dmErr != nil || !ok { return nil, fmt.Errorf("%w: cannot edit this message", ErrForbidden) } @@ -270,12 +271,12 @@ func (s *MessageService) EditMessage(userID, msgID int64, rawContent string) (*E } // EditMessage checks ownership internally. - if err := s.db.EditMessage(msgID, userID, content); err != nil { + if err := s.st.EditMessage(msgID, userID, content); err != nil { return nil, fmt.Errorf("%w: cannot edit this message", ErrForbidden) } // Re-fetch for updated edited_at timestamp. - msg, err = s.db.GetMessage(msgID) + msg, err = s.st.GetMessage(msgID) if err != nil || msg == nil { slog.Error("MessageService.EditMessage GetMessage after edit", "err", err, "msg_id", msgID) return nil, fmt.Errorf("%w: edit saved but broadcast failed", ErrInternal) @@ -295,7 +296,7 @@ func (s *MessageService) EditMessage(userID, msgID int64, rawContent string) (*E } if isDM { - participantIDs, pErr := s.db.GetDMParticipantIDs(msg.ChannelID) + participantIDs, pErr := s.st.GetDMParticipantIDs(msg.ChannelID) if pErr != nil { slog.Error("MessageService.EditMessage GetDMParticipantIDs", "err", pErr, "channel_id", msg.ChannelID) } else { @@ -319,16 +320,16 @@ func (s *MessageService) DeleteMessage(userID, msgID int64) (*DeleteMessageResul return nil, fmt.Errorf("%w: message_id must be positive integer", ErrBadRequest) } - msg, err := s.db.GetMessage(msgID) + msg, err := s.st.GetMessage(msgID) if err != nil || msg == nil { return nil, fmt.Errorf("%w: cannot delete this message", ErrForbidden) } - ch, chErr := s.db.GetChannel(msg.ChannelID) + ch, chErr := s.st.GetChannel(msg.ChannelID) isDM := chErr == nil && ch != nil && ch.Type == "dm" if isDM { - ok, dmErr := s.db.IsDMParticipant(userID, msg.ChannelID) + ok, dmErr := s.st.IsDMParticipant(userID, msg.ChannelID) if dmErr != nil || !ok { return nil, fmt.Errorf("%w: cannot delete this message", ErrForbidden) } @@ -342,12 +343,12 @@ func (s *MessageService) DeleteMessage(userID, msgID int64) (*DeleteMessageResul } isMod := !isDM && s.perms.HasChannelPerm(userID, msg.ChannelID, permissions.ManageMessages) - if err := s.db.DeleteMessage(msgID, userID, isMod); err != nil { + if err := s.st.DeleteMessage(msgID, userID, isMod); err != nil { return nil, fmt.Errorf("%w: cannot delete this message", ErrForbidden) } slog.Debug("message deleted", "user_id", userID, "msg_id", msgID, "channel_id", msg.ChannelID, "is_mod", isMod) - _ = s.db.LogAudit(userID, "message_delete", "message", msgID, + _ = s.st.LogAudit(userID, "message_delete", "message", msgID, fmt.Sprintf("channel %d, mod_action=%v", msg.ChannelID, isMod)) result := &DeleteMessageResult{ @@ -358,7 +359,7 @@ func (s *MessageService) DeleteMessage(userID, msgID int64) (*DeleteMessageResul } if isDM { - participantIDs, pErr := s.db.GetDMParticipantIDs(msg.ChannelID) + participantIDs, pErr := s.st.GetDMParticipantIDs(msg.ChannelID) if pErr != nil { slog.Error("MessageService.DeleteMessage GetDMParticipantIDs", "err", pErr, "channel_id", msg.ChannelID) } else { @@ -403,7 +404,7 @@ func (s *MessageService) handleReaction(userID, msgID int64, emoji string, add b return nil, fmt.Errorf("%w: emoji contains unsafe content", ErrBadRequest) } - msg, err := s.db.GetMessage(msgID) + msg, err := s.st.GetMessage(msgID) if err != nil || msg == nil { return nil, fmt.Errorf("%w: message not found", ErrForbidden) } @@ -411,11 +412,11 @@ func (s *MessageService) handleReaction(userID, msgID int64, emoji string, add b return nil, fmt.Errorf("%w: cannot react to deleted message", ErrDeletedMessage) } - ch, chErr := s.db.GetChannel(msg.ChannelID) + ch, chErr := s.st.GetChannel(msg.ChannelID) isDM := chErr == nil && ch != nil && ch.Type == "dm" if isDM { - ok, dmErr := s.db.IsDMParticipant(userID, msg.ChannelID) + ok, dmErr := s.st.IsDMParticipant(userID, msg.ChannelID) if dmErr != nil || !ok { return nil, fmt.Errorf("%w: not a DM participant", ErrForbidden) } @@ -427,13 +428,13 @@ func (s *MessageService) handleReaction(userID, msgID int64, emoji string, add b action := "add" if add { - if err := s.db.AddReaction(msgID, userID, emoji); err != nil { + if err := s.st.AddReaction(msgID, userID, emoji); err != nil { slog.Warn("MessageService.AddReaction", "err", err, "msg_id", msgID, "user_id", userID) return nil, fmt.Errorf("%w: reaction already exists", ErrConflict) } } else { action = "remove" - if err := s.db.RemoveReaction(msgID, userID, emoji); err != nil { + if err := s.st.RemoveReaction(msgID, userID, emoji); err != nil { slog.Warn("MessageService.RemoveReaction", "err", err, "msg_id", msgID, "user_id", userID) return nil, fmt.Errorf("%w: reaction not found", ErrBadRequest) } @@ -449,7 +450,7 @@ func (s *MessageService) handleReaction(userID, msgID int64, emoji string, add b } if isDM { - participantIDs, pErr := s.db.GetDMParticipantIDs(msg.ChannelID) + participantIDs, pErr := s.st.GetDMParticipantIDs(msg.ChannelID) if pErr != nil { slog.Error("MessageService.handleReaction GetDMParticipantIDs", "err", pErr, "channel_id", msg.ChannelID) } else { @@ -466,14 +467,14 @@ func (s *MessageService) GetMessages(userID, channelID, before int64, limit int) return nil, false, fmt.Errorf("%w: channel_id must be positive", ErrBadRequest) } - ch, err := s.db.GetChannel(channelID) + ch, err := s.st.GetChannel(channelID) if err != nil || ch == nil { return nil, false, fmt.Errorf("%w: channel not found", ErrNotFound) } // Permission check. if ch.Type == "dm" { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil || !ok { return nil, false, fmt.Errorf("%w: access denied", ErrForbidden) } @@ -491,7 +492,7 @@ func (s *MessageService) GetMessages(userID, channelID, before int64, limit int) } // Fetch one extra to detect has_more. - msgs, err := s.db.GetMessagesForAPI(channelID, before, limit+1, userID) + msgs, err := s.st.GetMessagesForAPI(channelID, before, limit+1, userID) if err != nil { slog.Error("MessageService.GetMessages", "err", err, "channel_id", channelID) return nil, false, fmt.Errorf("%w: failed to fetch messages", ErrInternal) @@ -519,19 +520,19 @@ func (s *MessageService) SearchMessages(userID int64, query string, channelID *i // Single-channel search. if channelID != nil && *channelID > 0 { - ch, err := s.db.GetChannel(*channelID) + ch, err := s.st.GetChannel(*channelID) if err != nil || ch == nil { return nil, fmt.Errorf("%w: channel not found", ErrNotFound) } if ch.Type == "dm" { - ok, err := s.db.IsDMParticipant(userID, *channelID) + ok, err := s.st.IsDMParticipant(userID, *channelID) if err != nil || !ok { return nil, fmt.Errorf("%w: access denied", ErrForbidden) } } else if !s.perms.HasChannelPerm(userID, *channelID, permissions.ReadMessages) { return nil, fmt.Errorf("%w: access denied", ErrForbidden) } - results, err := s.db.SearchMessages(query, channelID, limit) + results, err := s.st.SearchMessages(query, channelID, limit) if err != nil { return nil, fmt.Errorf("%w: search failed", ErrInternal) } @@ -547,7 +548,7 @@ func (s *MessageService) SearchMessages(userID int64, query string, channelID *i return nil, nil } - results, err := s.db.SearchMessagesInChannels(query, accessibleIDs, limit) + results, err := s.st.SearchMessagesInChannels(query, accessibleIDs, limit) if err != nil { return nil, fmt.Errorf("%w: search failed", ErrInternal) } @@ -559,19 +560,19 @@ func (s *MessageService) GetPinnedMessages(userID, channelID int64) ([]db.Messag if channelID <= 0 { return nil, fmt.Errorf("%w: channel_id must be positive", ErrBadRequest) } - ch, err := s.db.GetChannel(channelID) + ch, err := s.st.GetChannel(channelID) if err != nil || ch == nil { return nil, fmt.Errorf("%w: channel not found", ErrNotFound) } if ch.Type == "dm" { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil || !ok { return nil, fmt.Errorf("%w: access denied", ErrForbidden) } } else if !s.perms.HasChannelPerm(userID, channelID, permissions.ReadMessages) { return nil, fmt.Errorf("%w: access denied", ErrForbidden) } - msgs, err := s.db.GetPinnedMessages(channelID, userID) + msgs, err := s.st.GetPinnedMessages(channelID, userID) if err != nil { return nil, fmt.Errorf("%w: failed to fetch pinned messages", ErrInternal) } @@ -583,12 +584,12 @@ func (s *MessageService) SetMessagePinned(userID, channelID, msgID int64, pinned if channelID <= 0 || msgID <= 0 { return fmt.Errorf("%w: invalid IDs", ErrBadRequest) } - ch, err := s.db.GetChannel(channelID) + ch, err := s.st.GetChannel(channelID) if err != nil || ch == nil { return fmt.Errorf("%w: channel not found", ErrNotFound) } if ch.Type == "dm" { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil || !ok { return fmt.Errorf("%w: access denied", ErrForbidden) } @@ -596,16 +597,16 @@ func (s *MessageService) SetMessagePinned(userID, channelID, msgID int64, pinned return fmt.Errorf("%w: missing MANAGE_MESSAGES permission", ErrForbidden) } // Verify message belongs to this channel. - msg, err := s.db.GetMessage(msgID) + msg, err := s.st.GetMessage(msgID) if err != nil || msg == nil || msg.ChannelID != channelID { return fmt.Errorf("%w: message not found in this channel", ErrNotFound) } - return s.db.SetMessagePinned(msgID, pinned) + return s.st.SetMessagePinned(msgID, pinned) } // GetAccessibleChannelIDs returns all channel IDs the user can read. func (s *MessageService) GetAccessibleChannelIDs(userID int64) ([]int64, error) { - channels, err := s.db.ListChannels() + channels, err := s.st.ListChannels() if err != nil { return nil, fmt.Errorf("%w: failed to list channels", ErrInternal) } @@ -618,7 +619,7 @@ func (s *MessageService) GetAccessibleChannelIDs(userID int64) ([]int64, error) isAdmin := permissions.HasAdmin(role.Permissions) var overrides map[int64]db.ChannelOverride if !isAdmin { - overrides, _ = s.db.GetAllChannelPermissionsForRole(role.ID) + overrides, _ = s.st.GetAllChannelPermissionsForRole(role.ID) if overrides == nil { overrides = make(map[int64]db.ChannelOverride) } @@ -641,7 +642,7 @@ func (s *MessageService) GetAccessibleChannelIDs(userID int64) ([]int64, error) } // Also include DM channels the user participates in. - dmChannels, err := s.db.GetUserDMChannels(userID) + dmChannels, err := s.st.GetUserDMChannels(userID) if err == nil { for _, dmc := range dmChannels { ids = append(ids, dmc.ChannelID) @@ -654,16 +655,16 @@ func (s *MessageService) GetAccessibleChannelIDs(userID int64) ([]int64, error) // checkSendPermission validates send permission for DM and non-DM channels. func (s *MessageService) checkSendPermission(userID, channelID int64, isDM bool) error { if isDM { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil { return fmt.Errorf("%w: failed to check DM participation", ErrInternal) } if !ok { return fmt.Errorf("%w: not a participant in this DM", ErrForbidden) } - recipient, err := s.db.GetDMRecipient(channelID, userID) + recipient, err := s.st.GetDMRecipient(channelID, userID) if err == nil && recipient != nil { - blocked, blkErr := s.db.IsEitherBlocked(userID, recipient.ID) + blocked, blkErr := s.st.IsEitherBlocked(userID, recipient.ID) if blkErr != nil { return fmt.Errorf("%w: failed to check block status", ErrInternal) } diff --git a/Server/service/permission.go b/Server/service/permission.go index e2132bef..018d3c22 100644 --- a/Server/service/permission.go +++ b/Server/service/permission.go @@ -6,6 +6,7 @@ import ( "github.com/owncord/server/db" "github.com/owncord/server/permissions" + "github.com/owncord/server/store" ) // cachedPerms holds a snapshot of a user's role and channel overrides. @@ -24,7 +25,7 @@ const permCacheTTL = 30 * time.Second // at scale. The cache is populated lazily on first access and invalidated // on role or channel override changes. type PermissionService struct { - db *db.DB + st store.Store checker *permissions.Checker mu sync.RWMutex @@ -32,9 +33,9 @@ type PermissionService struct { } // NewPermissionService creates a PermissionService backed by the given DB. -func NewPermissionService(database *db.DB, checker *permissions.Checker) *PermissionService { +func NewPermissionService(st store.Store, checker *permissions.Checker) *PermissionService { return &PermissionService{ - db: database, + st: st, checker: checker, cache: make(map[int64]*cachedPerms), } @@ -60,7 +61,7 @@ func (s *PermissionService) HasChannelPerm(userID, channelID, perm int64) bool { // For regular channels it uses cached role-based permission checks. func (s *PermissionService) RequireChannelAccess(userID int64, channelType string, channelID, perm int64) error { if channelType == "dm" { - ok, err := s.db.IsDMParticipant(userID, channelID) + ok, err := s.st.IsDMParticipant(userID, channelID) if err != nil { return err } @@ -80,9 +81,9 @@ func (s *PermissionService) GetRoleForUser(userID int64) (*db.Role, error) { cp := s.getOrPopulate(userID) if cp == nil { // Cache miss, fall back to direct DB query. - return s.db.GetRoleForUser(userID) + return s.st.GetRoleForUser(userID) } - return s.db.GetRoleByID(cp.roleID) + return s.st.GetRoleByID(cp.roleID) } // InvalidateUser removes cached permissions for a specific user. @@ -127,11 +128,11 @@ func (s *PermissionService) getOrPopulate(userID int64) *cachedPerms { s.mu.RUnlock() // Populate. - role, err := s.db.GetRoleForUser(userID) + role, err := s.st.GetRoleForUser(userID) if err != nil || role == nil { return nil } - overrides, err := s.db.GetAllChannelPermissionsForRole(role.ID) + overrides, err := s.st.GetAllChannelPermissionsForRole(role.ID) if err != nil { // Fall back to uncached if override fetch fails. overrides = make(map[int64]db.ChannelOverride) diff --git a/Server/service/service.go b/Server/service/service.go index 74e839a5..6358bcc4 100644 --- a/Server/service/service.go +++ b/Server/service/service.go @@ -6,8 +6,8 @@ package service import ( "github.com/owncord/server/auth" - "github.com/owncord/server/db" "github.com/owncord/server/permissions" + "github.com/owncord/server/store" ) // Services bundles all domain services for dependency injection. @@ -16,15 +16,23 @@ type Services struct { Messages *MessageService Channels *ChannelService Permissions *PermissionService + Users *UserService + DMs *DMService + Invites *InviteService + Blocks *BlockService } // New creates all domain services wired together. -func New(database *db.DB, limiter *auth.RateLimiter) *Services { - permChecker := permissions.NewChecker(database) - permSvc := NewPermissionService(database, permChecker) +func New(st store.Store, limiter *auth.RateLimiter) *Services { + permChecker := permissions.NewChecker(st) + permSvc := NewPermissionService(st, permChecker) return &Services{ - Messages: NewMessageService(database, permSvc, limiter), - Channels: NewChannelService(database, permSvc), + Messages: NewMessageService(st, permSvc, limiter), + Channels: NewChannelService(st, permSvc), Permissions: permSvc, + Users: NewUserService(st), + DMs: NewDMService(st), + Invites: NewInviteService(st), + Blocks: NewBlockService(st), } } diff --git a/Server/service/user.go b/Server/service/user.go new file mode 100644 index 00000000..e5d80a16 --- /dev/null +++ b/Server/service/user.go @@ -0,0 +1,69 @@ +package service + +import ( + "fmt" + "log/slog" + + "github.com/owncord/server/db" + "github.com/owncord/server/store" +) + +// UserService handles user profile and session operations. +type UserService struct { + st store.Store +} + +// NewUserService creates a UserService. +func NewUserService(st store.Store) *UserService { + return &UserService{st: st} +} + +// UpdateProfile updates a user's username and/or avatar. +// Returns the updated user for response building. +func (s *UserService) UpdateProfile(userID int64, username string, avatar *string) (*db.User, error) { + if err := s.st.UpdateUserProfile(userID, username, avatar); err != nil { + return nil, fmt.Errorf("%w: %v", ErrInternal, err) + } + user, err := s.st.GetUserByID(userID) + if err != nil { + return nil, fmt.Errorf("%w: failed to fetch updated user", ErrInternal) + } + _ = s.st.LogAudit(userID, "profile_update", "user", userID, + fmt.Sprintf("username=%s", username)) + slog.Info("profile updated", "user_id", userID, "username", username) + return user, nil +} + +// ChangePassword verifies the old password hash matches, then updates. +// Returns the number of other sessions revoked. +func (s *UserService) ChangePassword(userID int64, newPasswordHash string, keepSessionID int64) (int64, error) { + if err := s.st.UpdateUserPassword(userID, newPasswordHash); err != nil { + return 0, fmt.Errorf("%w: failed to update password", ErrInternal) + } + revoked, err := s.st.DeleteOtherSessions(userID, keepSessionID) + if err != nil { + slog.Error("UserService.ChangePassword DeleteOtherSessions", "err", err, "user_id", userID) + } + _ = s.st.LogAudit(userID, "password_change", "user", userID, "password changed") + slog.Info("password changed", "user_id", userID, "sessions_revoked", revoked) + return revoked, nil +} + +// ListSessions returns all active sessions for a user. +func (s *UserService) ListSessions(userID int64) ([]db.Session, error) { + sessions, err := s.st.ListUserSessions(userID) + if err != nil { + return nil, fmt.Errorf("%w: failed to list sessions", ErrInternal) + } + return sessions, nil +} + +// RevokeSession deletes a specific session owned by the user. +func (s *UserService) RevokeSession(userID, sessionID int64) error { + if err := s.st.DeleteSessionByID(sessionID, userID); err != nil { + return fmt.Errorf("%w: session not found", ErrNotFound) + } + _ = s.st.LogAudit(userID, "session_revoke", "session", sessionID, "session revoked") + slog.Info("session revoked", "user_id", userID, "session_id", sessionID) + return nil +} diff --git a/Server/ws/emit_test.go b/Server/ws/emit_test.go index a6f338ae..935e844f 100644 --- a/Server/ws/emit_test.go +++ b/Server/ws/emit_test.go @@ -31,6 +31,7 @@ func newEmitTestHub() *Hub { register: make(chan *Client, 16), unregister: make(chan *Client, 16), stop: make(chan struct{}), + pubsub: NewPubSub(), replayBuf: NewEventRingBuffer(100), voiceKeyHolders: make(map[int64]int64), } @@ -42,6 +43,12 @@ func registerEmitTestClient(h *Hub, userID, channelID int64) chan []byte { send := make(chan []byte, 64) c := NewTestClientWithChannel(h, userID, channelID, send) h.clients[userID] = c + // Subscribe to pub/sub topics so deliverBroadcast can reach this client. + h.pubsub.Subscribe(c, TopicGlobal) + h.pubsub.Subscribe(c, UserTopic(userID)) + if channelID > 0 { + h.pubsub.Subscribe(c, ChannelTopic(channelID)) + } return send } @@ -51,6 +58,15 @@ func registerEmitTestVoiceClient(h *Hub, userID, channelID, voiceChID int64) cha c := NewTestClientWithChannel(h, userID, channelID, send) SetClientVoiceChID(c, voiceChID) h.clients[userID] = c + // Subscribe to pub/sub topics so deliverBroadcast can reach this client. + h.pubsub.Subscribe(c, TopicGlobal) + h.pubsub.Subscribe(c, UserTopic(userID)) + if channelID > 0 { + h.pubsub.Subscribe(c, ChannelTopic(channelID)) + } + if voiceChID > 0 { + h.pubsub.Subscribe(c, ChannelTopic(voiceChID)) + } return send } diff --git a/Server/ws/handlers.go b/Server/ws/handlers.go index 559da71d..dc4a3ec9 100644 --- a/Server/ws/handlers.go +++ b/Server/ws/handlers.go @@ -232,17 +232,13 @@ func (h *Hub) requireChannelPerm(c *Client, channelID int64, perm int64, permLab // This is correct for typing indicators but would be incorrect for messages // that should survive reconnection replay. func (h *Hub) broadcastExclude(channelID, excludeUserID int64, msg []byte) { - h.mu.RLock() - defer h.mu.RUnlock() - for uid, c := range h.clients { - if uid == excludeUserID { - continue - } - if channelID != 0 && c.getChannelID() != channelID { - continue - } - c.sendMsg(msg) + if channelID == 0 { + // Global broadcast excluding one user — use the global topic. + h.pubsub.Publish(TopicGlobal, msg, excludeUserID) + return } + // Channel-scoped broadcast excluding one user. + h.pubsub.Publish(ChannelTopic(channelID), msg, excludeUserID) } // broadcastToDMParticipants sends a message to all participants of a DM channel