diff --git a/db/map_mods.go b/db/map_mods.go index 48c7356..d6819dd 100644 --- a/db/map_mods.go +++ b/db/map_mods.go @@ -7,6 +7,7 @@ import ( type MapModStatus string type MapModType string +type MapModSort string const ( ModStatusPending MapModStatus = "Pending" @@ -16,6 +17,10 @@ const ( ModTypeIssue MapModType = "Issue" ModTypeSuggestion MapModType = "Suggestion" + + ModSortRecent MapModSort = "recent" + ModSortStatus MapModSort = "status" + ModSortType MapModSort = "type" ) type MapMod struct { @@ -42,15 +47,14 @@ func (mod *MapMod) AfterFind(*gorm.DB) (err error) { } // GetMapMods retrieves a page of map mods and all replies belonging to those mods. -func GetMapMods(id int, page int, limit int) ([]*MapMod, error) { +func GetMapMods(id int, page int, limit int, statuses []MapModStatus, modType *MapModType, sort MapModSort) ([]*MapMod, error) { var mods = make([]*MapMod, 0) - result := SQL. + result := mapModsQuery(id, statuses, modType). Joins("Author"). Preload("Replies"). Preload("Replies.Author"). - Where("map_mods.map_id = ?", id). - Order("map_mods.id ASC"). + Order(mapModsOrder(sort)). Limit(limit). Offset(page * limit). Find(&mods) @@ -63,12 +67,10 @@ func GetMapMods(id int, page int, limit int) ([]*MapMod, error) { } // GetMapModsCount gets the total number of mods for a map. -func GetMapModsCount(id int) (int64, error) { +func GetMapModsCount(id int, statuses []MapModStatus, modType *MapModType) (int64, error) { var count int64 - result := SQL. - Model(&MapMod{}). - Where("map_id = ?", id). + result := mapModsQuery(id, statuses, modType). Count(&count) if result.Error != nil { @@ -78,6 +80,31 @@ func GetMapModsCount(id int) (int64, error) { return count, nil } +func mapModsQuery(id int, statuses []MapModStatus, modType *MapModType) *gorm.DB { + query := SQL.Model(&MapMod{}).Where("map_mods.map_id = ?", id) + + if len(statuses) > 0 { + query = query.Where("map_mods.status IN ?", statuses) + } + + if modType != nil { + query = query.Where("map_mods.type = ?", *modType) + } + + return query +} + +func mapModsOrder(sort MapModSort) string { + switch sort { + case ModSortStatus: + return "CASE map_mods.status WHEN 'Pending' THEN 0 WHEN 'Accepted' THEN 1 WHEN 'Denied' THEN 2 WHEN 'Ignored' THEN 3 ELSE 4 END ASC, map_mods.timestamp DESC, map_mods.id DESC" + case ModSortType: + return "CASE map_mods.type WHEN 'None' THEN 0 WHEN 'Issue' THEN 1 WHEN 'Suggestion' THEN 2 ELSE 3 END ASC, map_mods.timestamp DESC, map_mods.id DESC" + default: + return "map_mods.timestamp DESC, map_mods.id DESC" + } +} + // GetModById Gets a mod by its id func GetModById(id int) (*MapMod, error) { var mod *MapMod diff --git a/handlers/map_mods.go b/handlers/map_mods.go index 96f22cf..c7d8608 100644 --- a/handlers/map_mods.go +++ b/handlers/map_mods.go @@ -1,6 +1,7 @@ package handlers import ( + "fmt" "github.com/Quaver/api2/db" "github.com/Quaver/api2/enums" "github.com/gin-gonic/gin" @@ -20,15 +21,19 @@ func GetMapMods(c *gin.Context) *APIError { return APIErrorBadRequest("Invalid id") } - page, limit := getMapModsPagination(c) + page, limit, statuses, modType, sort, paginationError := getMapModsPagination(c) - mods, err := db.GetMapMods(id, page, limit) + if paginationError != nil { + return paginationError + } + + mods, err := db.GetMapMods(id, page, limit, statuses, modType, sort) if err != nil { return APIErrorServerError("Error retrieving map mods from db", err) } - total, err := db.GetMapModsCount(id) + total, err := db.GetMapModsCount(id, statuses, modType) if err != nil { return APIErrorServerError("Error retrieving map mods count from db", err) @@ -38,8 +43,78 @@ func GetMapMods(c *gin.Context) *APIError { return nil } -func getMapModsPagination(c *gin.Context) (int, int) { - return getQueryPage(c), getQueryLimit(c, defaultMapModLimit) +func getMapModsPagination(c *gin.Context) (int, int, []db.MapModStatus, *db.MapModType, db.MapModSort, *APIError) { + statuses, err := getMapModStatuses(c.Query("status")) + + if err != nil { + return 0, 0, nil, nil, db.ModSortRecent, APIErrorBadRequest("Invalid mod status filter") + } + + modType, err := getMapModType(c.Query("type")) + + if err != nil { + return 0, 0, nil, nil, db.ModSortRecent, APIErrorBadRequest("Invalid mod type filter") + } + + sort, err := getMapModSort(c.Query("sort")) + + if err != nil { + return 0, 0, nil, nil, db.ModSortRecent, APIErrorBadRequest("Invalid mod sort") + } + + return getQueryPage(c), getQueryLimit(c, defaultMapModLimit), statuses, modType, sort, nil +} + +func getMapModStatuses(value string) ([]db.MapModStatus, error) { + if value == "" { + return nil, nil + } + + statuses := make([]db.MapModStatus, 0) + seen := make(map[db.MapModStatus]struct{}) + + for _, value := range strings.Split(value, ",") { + status := db.MapModStatus(strings.TrimSpace(value)) + + if status != db.ModStatusPending && status != db.ModStatusAccepted && status != db.ModStatusDenied && status != db.ModStatusIgnored { + return nil, fmt.Errorf("invalid map mod status") + } + + if _, exists := seen[status]; !exists { + statuses = append(statuses, status) + seen[status] = struct{}{} + } + } + + return statuses, nil +} + +func getMapModType(value string) (*db.MapModType, error) { + if value == "" { + return nil, nil + } + + modType := db.MapModType(value) + + if modType != db.ModTypeIssue && modType != db.ModTypeSuggestion { + return nil, fmt.Errorf("invalid map mod type") + } + + return &modType, nil +} + +func getMapModSort(value string) (db.MapModSort, error) { + if value == "" { + return db.ModSortRecent, nil + } + + sort := db.MapModSort(value) + + if sort != db.ModSortRecent && sort != db.ModSortStatus && sort != db.ModSortType { + return db.ModSortRecent, fmt.Errorf("invalid map mod sort") + } + + return sort, nil } // GetMapMod gets a single mod for a map, including all of its replies.