diff --git a/Gopkg.lock b/Gopkg.lock index b63698a555..9cfd7bbb58 100644 --- a/Gopkg.lock +++ b/Gopkg.lock @@ -1,6 +1,14 @@ # This file is autogenerated, do not edit; changes may be undone by the next 'dep ensure'. +[[projects]] + digest = "1:14a90eb1290bd0aa42848afa9ee2c9ce247abdb63c2cf614331670119760962f" + name = "github.com/Masterminds/squirrel" + packages = ["."] + pruneopts = "UT" + revision = "fa735ea14f09f8685fcf5e75db091feb5a410730" + version = "v1.1" + [[projects]] digest = "1:0803645e1f57fb5271a6edc7570b9ea59bac2e5de67957075a43f3d74c8dbd97" name = "github.com/NYTimes/gziphandler" @@ -375,6 +383,22 @@ revision = "5c8c8bd35d3832f5d134ae1e1e375b69a4d25242" version = "v1.0.1" +[[projects]] + branch = "master" + digest = "1:a7fc52742a5d011497b6a24024c857d260f809083424cd84110c9e40c34f64fc" + name = "github.com/lann/builder" + packages = ["."] + pruneopts = "UT" + revision = "47ae307949d02aa1f1069fdafc00ca08e1dbabac" + +[[projects]] + branch = "master" + digest = "1:225499d25a9f1486f3b77cdc4f7d6590c506c3839bb9d8497113f6d19676d54a" + name = "github.com/lann/ps" + packages = ["."] + pruneopts = "UT" + revision = "62de8c46ede02a7675c4c79c84883eb164cb71e3" + [[projects]] digest = "1:8ef506fc2bb9ced9b151dafa592d4046063d744c646c1bbe801982ce87e4bc24" name = "github.com/lib/pq" @@ -1061,6 +1085,7 @@ analyzer-name = "dep" analyzer-version = 1 input-imports = [ + "github.com/Masterminds/squirrel", "github.com/NYTimes/gziphandler", "github.com/avct/uasurfer", "github.com/blang/semver", diff --git a/app/channel_test.go b/app/channel_test.go index 273ce1f664..a0286b22ec 100644 --- a/app/channel_test.go +++ b/app/channel_test.go @@ -5,6 +5,7 @@ package app import ( "fmt" + "sort" "strings" "testing" @@ -305,6 +306,9 @@ func TestCreateGroupChannelCreatesChannelMemberHistoryRecord(t *testing.T) { assert.Equal(t, channel.Id, history.ChannelId) channelMemberHistoryUserIds = append(channelMemberHistoryUserIds, history.UserId) } + + sort.Strings(groupUserIds) + sort.Strings(channelMemberHistoryUserIds) assert.Equal(t, groupUserIds, channelMemberHistoryUserIds) } } diff --git a/i18n/en.json b/i18n/en.json index 1639e9b900..000c257e18 100644 --- a/i18n/en.json +++ b/i18n/en.json @@ -6286,6 +6286,10 @@ "id": "store.sql_team.update_last_team_icon_update.app_error", "translation": "Unable to update the date of the last team icon update" }, + { + "id": "store.sql_user.app_error", + "translation": "Failed to build query" + }, { "id": "store.sql_user.analytics_daily_active_users.app_error", "translation": "Unable to get the active users during the requested period" diff --git a/store/sqlstore/user_store.go b/store/sqlstore/user_store.go index 964266faf4..7c97cfa16b 100644 --- a/store/sqlstore/user_store.go +++ b/store/sqlstore/user_store.go @@ -7,9 +7,9 @@ import ( "database/sql" "fmt" "net/http" - "strconv" "strings" + sq "github.com/Masterminds/squirrel" "github.com/mattermost/gorp" "github.com/mattermost/mattermost-server/einterfaces" @@ -35,6 +35,9 @@ var ( type SqlUserStore struct { SqlStore metrics einterfaces.MetricsInterface + + // usersQuery is a starting point for all queries that return one or more Users. + usersQuery sq.SelectBuilder } var profilesInChannelCache *utils.Cache = utils.NewLru(PROFILES_IN_CHANNEL_CACHE_SIZE) @@ -64,6 +67,14 @@ func NewSqlUserStore(sqlStore SqlStore, metrics einterfaces.MetricsInterface) st metrics: metrics, } + us.usersQuery = sq. + Select("u.*"). + From("Users u") + + if us.DriverName() == model.DATABASE_DRIVER_POSTGRES { + us.usersQuery = us.usersQuery.PlaceholderFormat(sq.Dollar) + } + for _, db := range sqlStore.GetAllConns() { table := db.AddTableWithName(model.User{}, "Users").SetKeys(false, "Id") table.ColMap("Id").SetMaxSize(26) @@ -205,7 +216,7 @@ func (us SqlUserStore) UpdateLastPictureUpdate(userId string) store.StoreChannel curTime := model.GetMillis() if _, err := us.GetMaster().Exec("UPDATE Users SET LastPictureUpdate = :Time, UpdateAt = :Time WHERE Id = :UserId", map[string]interface{}{"Time": curTime, "UserId": userId}); err != nil { - result.Err = model.NewAppError("SqlUserStore.UpdateUpdateAt", "store.sql_user.update_last_picture_update.app_error", nil, "user_id="+userId, http.StatusInternalServerError) + result.Err = model.NewAppError("SqlUserStore.UpdateLastPictureUpdate", "store.sql_user.update_last_picture_update.app_error", nil, "user_id="+userId, http.StatusInternalServerError) } else { result.Data = userId } @@ -215,7 +226,7 @@ func (us SqlUserStore) UpdateLastPictureUpdate(userId string) store.StoreChannel func (us SqlUserStore) ResetLastPictureUpdate(userId string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { if _, err := us.GetMaster().Exec("UPDATE Users SET LastPictureUpdate = :Time, UpdateAt = :Time WHERE Id = :UserId", map[string]interface{}{"Time": 0, "UserId": userId}); err != nil { - result.Err = model.NewAppError("SqlUserStore.UpdateUpdateAt", "store.sql_user.update_last_picture_update.app_error", nil, "user_id="+userId, http.StatusInternalServerError) + result.Err = model.NewAppError("SqlUserStore.ResetLastPictureUpdate", "store.sql_user.update_last_picture_update.app_error", nil, "user_id="+userId, http.StatusInternalServerError) } else { result.Data = userId } @@ -228,9 +239,10 @@ func (us SqlUserStore) UpdateUpdateAt(userId string) store.StoreChannel { if _, err := us.GetMaster().Exec("UPDATE Users SET UpdateAt = :Time WHERE Id = :UserId", map[string]interface{}{"Time": curTime, "UserId": userId}); err != nil { result.Err = model.NewAppError("SqlUserStore.UpdateUpdateAt", "store.sql_user.update_update.app_error", nil, "user_id="+userId, http.StatusInternalServerError) - } else { - result.Data = userId + return } + + result.Data = curTime }) } @@ -321,21 +333,41 @@ func (us SqlUserStore) UpdateMfaActive(userId string, active bool) store.StoreCh func (us SqlUserStore) Get(id string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - if obj, err := us.GetReplica().Get(model.User{}, id); err != nil { - result.Err = model.NewAppError("SqlUserStore.Get", "store.sql_user.get.app_error", nil, "user_id="+id+", "+err.Error(), http.StatusInternalServerError) - } else if obj == nil { - result.Err = model.NewAppError("SqlUserStore.Get", store.MISSING_ACCOUNT_ERROR, nil, "user_id="+id, http.StatusNotFound) - } else { - result.Data = obj.(*model.User) + query := us.usersQuery.Where("Id = ?", id) + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.Get", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + user := &model.User{} + if err := us.GetReplica().SelectOne(user, queryString, args...); err == sql.ErrNoRows { + result.Err = model.NewAppError("SqlUserStore.Get", store.MISSING_ACCOUNT_ERROR, nil, "user_id="+id, http.StatusNotFound) + return + } else if err != nil { + result.Err = model.NewAppError("SqlUserStore.Get", "store.sql_user.get.app_error", nil, "user_id="+id+", "+err.Error(), http.StatusInternalServerError) + return + } + + result.Data = user }) } func (us SqlUserStore) GetAll() store.StoreChannel { return store.Do(func(result *store.StoreResult) { + query := us.usersQuery.OrderBy("Username ASC") + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAll", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + var data []*model.User - if _, err := us.GetReplica().Select(&data, "SELECT * FROM Users"); err != nil { + if _, err := us.GetReplica().Select(&data, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetAll", "store.sql_user.get.app_error", nil, err.Error(), http.StatusInternalServerError) + return } result.Data = data @@ -344,8 +376,19 @@ func (us SqlUserStore) GetAll() store.StoreChannel { func (us SqlUserStore) GetAllAfter(limit int, afterId string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { + query := us.usersQuery. + Where("Id > ?", afterId). + OrderBy("Id ASC"). + Limit(uint64(limit)) + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAllAfter", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + var data []*model.User - if _, err := us.GetReplica().Select(&data, "SELECT * FROM Users WHERE Id > :AfterId ORDER BY Id LIMIT :Limit", map[string]interface{}{"AfterId": afterId, "Limit": limit}); err != nil { + if _, err := us.GetReplica().Select(&data, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetAllAfter", "store.sql_user.get.app_error", nil, err.Error(), http.StatusInternalServerError) } @@ -367,61 +410,47 @@ func (s SqlUserStore) GetEtagForAllProfiles() store.StoreChannel { func (us SqlUserStore) GetAllProfiles(options *model.UserGetOptions) store.StoreChannel { isPostgreSQL := us.DriverName() == model.DATABASE_DRIVER_POSTGRES return store.Do(func(result *store.StoreResult) { - var users []*model.User - offset := options.Page * options.PerPage - limit := options.PerPage + query := us.usersQuery. + OrderBy("u.Username ASC"). + Offset(uint64(options.Page * options.PerPage)).Limit(uint64(options.PerPage)) - searchQuery := ` - SELECT * FROM Users - WHERE_CONDITION - ORDER BY Username ASC LIMIT :Limit OFFSET :Offset - ` + query = applyRoleFilter(query, options.Role, isPostgreSQL) - parameters := map[string]interface{}{"Offset": offset, "Limit": limit} - searchQuery = substituteWhereClause(searchQuery, options, parameters, isPostgreSQL) - - if _, err := us.GetReplica().Select(&users, searchQuery, parameters); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetAllProfiles", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + if options.Inactive { + query = query.Where("u.DeleteAt != 0") } + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAllProfiles", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAllProfiles", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } -func substituteWhereClause(searchQuery string, options *model.UserGetOptions, parameters map[string]interface{}, isPostgreSQL bool) string { - whereClause := "" - whereClauses := []string{} - if options.Role != "" { - whereClauses = append(whereClauses, getRoleFilter(isPostgreSQL)) - parameters["Role"] = fmt.Sprintf("%%%s%%", options.Role) - } - if options.Inactive { - whereClauses = append(whereClauses, " Users.DeleteAt != 0 ") +func applyRoleFilter(query sq.SelectBuilder, role string, isPostgreSQL bool) sq.SelectBuilder { + if role == "" { + return query } - if len(whereClauses) > 0 { - whereClause = strings.Join(whereClauses, " AND ") - searchQuery = strings.Replace(searchQuery, "WHERE_CONDITION", fmt.Sprintf(" WHERE %s ", whereClause), 1) - searchQuery = strings.Replace(searchQuery, "SEARCH_CLAUSE", fmt.Sprintf(" AND %s ", whereClause), 1) - } else { - searchQuery = strings.Replace(searchQuery, "WHERE_CONDITION", "", 1) - searchQuery = strings.Replace(searchQuery, "SEARCH_CLAUSE", "", 1) - } - - return searchQuery -} - -func getRoleFilter(isPostgreSQL bool) string { + roleParam := fmt.Sprintf("%%%s%%", role) if isPostgreSQL { - return fmt.Sprintf("Users.Roles like lower(%s)", ":Role") - } else { - return fmt.Sprintf("Users.Roles LIKE %s escape '*' ", ":Role") + return query.Where("u.Roles LIKE LOWER(?)", roleParam) } + + return query.Where("u.Roles LIKE ? ESCAPE '*'", roleParam) } func (s SqlUserStore) GetEtagForProfiles(teamId string) store.StoreChannel { @@ -437,33 +466,35 @@ func (s SqlUserStore) GetEtagForProfiles(teamId string) store.StoreChannel { func (us SqlUserStore) GetProfiles(options *model.UserGetOptions) store.StoreChannel { isPostgreSQL := us.DriverName() == model.DATABASE_DRIVER_POSTGRES - teamId := options.InTeamId - offset := options.Page * options.PerPage - limit := options.PerPage - - searchQuery := ` - SELECT Users.* FROM Users, TeamMembers - WHERE TeamMembers.TeamId = :TeamId AND Users.Id = TeamMembers.UserId AND TeamMembers.DeleteAt = 0 - SEARCH_CLAUSE - ORDER BY Users.Username ASC LIMIT :Limit OFFSET :Offset - ` - - parameters := map[string]interface{}{"TeamId": teamId, "Offset": offset, "Limit": limit} - searchQuery = substituteWhereClause(searchQuery, options, parameters, isPostgreSQL) - return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + Join("TeamMembers tm ON ( tm.UserId = u.Id AND tm.DeleteAt = 0 )"). + Where("tm.TeamId = ?", options.InTeamId). + OrderBy("u.Username ASC"). + Offset(uint64(options.Page * options.PerPage)).Limit(uint64(options.PerPage)) - if _, err := us.GetReplica().Select(&users, searchQuery, parameters); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetProfiles", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { + query = applyRoleFilter(query, options.Role, isPostgreSQL) - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + if options.Inactive { + query = query.Where("u.DeleteAt != 0") } + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfiles", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfiles", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + result.Data = users }) } @@ -492,67 +523,66 @@ func (us SqlUserStore) InvalidateProfilesInChannelCache(channelId string) { func (us SqlUserStore) GetProfilesInChannel(channelId string, offset int, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + Join("ChannelMembers cm ON ( cm.UserId = u.Id )"). + Where("cm.ChannelId = ?", channelId). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - query := ` - SELECT - Users.* - FROM - Users, ChannelMembers - WHERE - ChannelMembers.ChannelId = :ChannelId - AND Users.Id = ChannelMembers.UserId - ORDER BY - Users.Username ASC - LIMIT :Limit OFFSET :Offset - ` - - if _, err := us.GetReplica().Select(&users, query, map[string]interface{}{"ChannelId": channelId, "Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetProfilesInChannel", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesInChannel", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesInChannel", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } func (us SqlUserStore) GetProfilesInChannelByStatus(channelId string, offset int, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User - - query := ` - SELECT - Users.* - FROM Users - INNER JOIN ChannelMembers ON Users.Id = ChannelMembers.UserId - LEFT JOIN Status ON Users.Id = Status.UserId - WHERE - ChannelMembers.ChannelId = :ChannelId - ORDER BY - CASE Status + query := us.usersQuery. + Join("ChannelMembers cm ON ( cm.UserId = u.Id )"). + LeftJoin("Status s ON ( s.UserId = u.Id )"). + Where("cm.ChannelId = ?", channelId). + OrderBy(` + CASE s.Status WHEN 'online' THEN 1 WHEN 'away' THEN 2 WHEN 'dnd' THEN 3 ELSE 4 - END, - Users.Username ASC - LIMIT :Limit OFFSET :Offset - ` + END + `). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - if _, err := us.GetReplica().Select(&users, query, map[string]interface{}{"ChannelId": channelId, "Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetProfilesInChannelByStatus", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesInChannelByStatus", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesInChannelByStatus", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } @@ -576,127 +606,130 @@ func (us SqlUserStore) GetAllProfilesInChannel(channelId string, allowFromCache } } + query := us.usersQuery. + Join("ChannelMembers cm ON ( cm.UserId = u.Id )"). + Where("cm.ChannelId = ?", channelId). + Where("u.DeleteAt = 0"). + OrderBy("u.Username ASC") + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAllProfilesInChannel", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + var users []*model.User - - query := "SELECT Users.* FROM Users, ChannelMembers WHERE ChannelMembers.ChannelId = :ChannelId AND Users.Id = ChannelMembers.UserId AND Users.DeleteAt = 0" - - if _, err := us.GetReplica().Select(&users, query, map[string]interface{}{"ChannelId": channelId}); err != nil { + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetAllProfilesInChannel", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { + return + } - userMap := make(map[string]*model.User) + userMap := make(map[string]*model.User) - for _, u := range users { - u.Sanitize(map[string]bool{}) - userMap[u.Id] = u - } + for _, u := range users { + u.Sanitize(map[string]bool{}) + userMap[u.Id] = u + } - result.Data = userMap + result.Data = userMap - if allowFromCache { - profilesInChannelCache.AddWithExpiresInSecs(channelId, userMap, PROFILES_IN_CHANNEL_CACHE_SEC) - } + if allowFromCache { + profilesInChannelCache.AddWithExpiresInSecs(channelId, userMap, PROFILES_IN_CHANNEL_CACHE_SEC) } }) } func (us SqlUserStore) GetProfilesNotInChannel(teamId string, channelId string, offset int, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + Join("TeamMembers tm ON ( tm.UserId = u.Id AND tm.DeleteAt = 0 AND tm.TeamId = ? )", teamId). + LeftJoin("ChannelMembers cm ON ( cm.UserId = u.Id AND cm.ChannelId = ? )", channelId). + Where("cm.UserId IS NULL"). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - if _, err := us.GetReplica().Select(&users, ` - SELECT - u.* - FROM Users u - INNER JOIN TeamMembers tm - ON tm.UserId = u.Id - AND tm.TeamId = :TeamId - AND tm.DeleteAt = 0 - LEFT JOIN ChannelMembers cm - ON cm.UserId = u.Id - AND cm.ChannelId = :ChannelId - WHERE cm.UserId IS NULL - ORDER BY u.Username ASC - LIMIT :Limit OFFSET :Offset - `, map[string]interface{}{"TeamId": teamId, "ChannelId": channelId, "Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetProfilesNotInChannel", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesNotInChannel", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesNotInChannel", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } func (us SqlUserStore) GetProfilesWithoutTeam(offset int, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + Where(`( + SELECT + COUNT(0) + FROM + TeamMembers + WHERE + TeamMembers.UserId = u.Id + AND TeamMembers.DeleteAt = 0 + ) = 0`). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - query := ` - SELECT - * - FROM - Users - WHERE - (SELECT - COUNT(0) - FROM - TeamMembers - WHERE - TeamMembers.UserId = Users.Id - AND TeamMembers.DeleteAt = 0) = 0 - ORDER BY - Username ASC - LIMIT - :Limit - OFFSET - :Offset` - - if _, err := us.GetReplica().Select(&users, query, map[string]interface{}{"Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetProfilesWithoutTeam", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesWithoutTeam", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesWithoutTeam", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } func (us SqlUserStore) GetProfilesByUsernames(usernames []string, teamId string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { + query := us.usersQuery + + if teamId != "" { + query = query.Join("TeamMembers tm ON (tm.UserId = u.Id AND tm.TeamId = ?)", teamId) + } + + query = query. + Where(map[string]interface{}{ + "Username": usernames, + }). + OrderBy("u.Username ASC") + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesByUsernames", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + var users []*model.User - props := make(map[string]interface{}) - idQuery := "" - - for index, usernames := range usernames { - if len(idQuery) > 0 { - idQuery += ", " - } - - props["username"+strconv.Itoa(index)] = usernames - idQuery += ":username" + strconv.Itoa(index) - } - - var query string - if teamId == "" { - query = `SELECT * FROM Users WHERE Username IN (` + idQuery + `)` - } else { - query = `SELECT Users.* FROM Users INNER JOIN TeamMembers ON - Users.Id = TeamMembers.UserId AND Users.Username IN (` + idQuery + `) AND TeamMembers.TeamId = :TeamId ` - props["TeamId"] = teamId - } - - if _, err := us.GetReplica().Select(&users, query, props); err != nil { + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetProfilesByUsernames", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - result.Data = users + return } + + result.Data = users }) } @@ -707,65 +740,70 @@ type UserWithLastActivityAt struct { func (us SqlUserStore) GetRecentlyActiveUsersForTeam(teamId string, offset, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*UserWithLastActivityAt + query := us.usersQuery. + Column("s.LastActivityAt"). + Join("TeamMembers tm ON (tm.UserId = u.Id AND tm.TeamId = ?)", teamId). + Join("Status s ON (s.UserId = u.Id)"). + OrderBy("s.LastActivityAt DESC"). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - if _, err := us.GetReplica().Select(&users, ` - SELECT - u.*, - s.LastActivityAt - FROM Users AS u - INNER JOIN TeamMembers AS t ON u.Id = t.UserId - INNER JOIN Status AS s ON s.UserId = t.UserId - WHERE t.TeamId = :TeamId - ORDER BY s.LastActivityAt DESC - LIMIT :Limit OFFSET :Offset - `, map[string]interface{}{"TeamId": teamId, "Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetRecentlyActiveUsers", "store.sql_user.get_recently_active_users.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - userList := []*model.User{} - - for _, userWithLastActivityAt := range users { - u := userWithLastActivityAt.User - u.Sanitize(map[string]bool{}) - u.LastActivityAt = userWithLastActivityAt.LastActivityAt - userList = append(userList, &u) - } - - result.Data = userList + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetRecentlyActiveUsers", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*UserWithLastActivityAt + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetRecentlyActiveUsers", "store.sql_user.get_recently_active_users.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + userList := []*model.User{} + + for _, userWithLastActivityAt := range users { + u := userWithLastActivityAt.User + u.Sanitize(map[string]bool{}) + u.LastActivityAt = userWithLastActivityAt.LastActivityAt + userList = append(userList, &u) + } + + result.Data = userList }) } func (us SqlUserStore) GetNewUsersForTeam(teamId string, offset, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + Join("TeamMembers tm ON (tm.UserId = u.Id AND tm.TeamId = ?)", teamId). + OrderBy("u.CreateAt DESC"). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - if _, err := us.GetReplica().Select(&users, ` - SELECT - u.* - FROM Users AS u - INNER JOIN TeamMembers AS t ON u.Id = t.UserId - WHERE t.TeamId = :TeamId - ORDER BY u.CreateAt DESC - LIMIT :Limit OFFSET :Offset - `, map[string]interface{}{"TeamId": teamId, "Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetNewUsersForTeam", "store.sql_user.get_new_users.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetNewUsersForTeam", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetNewUsersForTeam", "store.sql_user.get_new_users.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } func (us SqlUserStore) GetProfileByIds(userIds []string, allowFromCache bool) store.StoreChannel { return store.Do(func(result *store.StoreResult) { users := []*model.User{} - props := make(map[string]interface{}) - idQuery := "" remainingUserIds := make([]string, 0) if allowFromCache { @@ -795,49 +833,61 @@ func (us SqlUserStore) GetProfileByIds(userIds []string, allowFromCache bool) st return } - for index, userId := range remainingUserIds { - if len(idQuery) > 0 { - idQuery += ", " - } + query := us.usersQuery. + Where(map[string]interface{}{ + "u.Id": remainingUserIds, + }). + OrderBy("u.Username ASC") - props["userId"+strconv.Itoa(index)] = userId - idQuery += ":userId" + strconv.Itoa(index) + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfileByIds", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } - if _, err := us.GetReplica().Select(&users, "SELECT * FROM Users WHERE Users.Id IN ("+idQuery+")", props); err != nil { + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetProfileByIds", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - - cpy := &model.User{} - *cpy = *u - profileByIdsCache.AddWithExpiresInSecs(cpy.Id, cpy, PROFILE_BY_IDS_CACHE_SEC) - } - - result.Data = users + return } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + + cpy := &model.User{} + *cpy = *u + profileByIdsCache.AddWithExpiresInSecs(cpy.Id, cpy, PROFILE_BY_IDS_CACHE_SEC) + } + + result.Data = users }) } func (us SqlUserStore) GetSystemAdminProfiles() store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + Where("Roles LIKE ?", "%system_admin%"). + OrderBy("u.Username ASC") - if _, err := us.GetReplica().Select(&users, "SELECT * FROM Users WHERE Roles LIKE :Roles", map[string]interface{}{"Roles": "%system_admin%"}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetSystemAdminProfiles", "store.sql_user.get_sysadmin_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - userMap := make(map[string]*model.User) - - for _, u := range users { - u.Sanitize(map[string]bool{}) - userMap[u.Id] = u - } - - result.Data = userMap + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetSystemAdminProfiles", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetSystemAdminProfiles", "store.sql_user.get_sysadmin_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + userMap := make(map[string]*model.User) + + for _, u := range users { + u.Sanitize(map[string]bool{}) + userMap[u.Id] = u + } + + result.Data = userMap }) } @@ -845,9 +895,16 @@ func (us SqlUserStore) GetByEmail(email string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { email = strings.ToLower(email) - user := model.User{} + query := us.usersQuery.Where("Email = ?", email) - if err := us.GetReplica().SelectOne(&user, "SELECT * FROM Users WHERE Email = :Email", map[string]interface{}{"Email": email}); err != nil { + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetByEmail", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + user := model.User{} + if err := us.GetReplica().SelectOne(&user, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetByEmail", store.MISSING_ACCOUNT_ERROR, nil, "email="+email+", "+err.Error(), http.StatusInternalServerError) } @@ -862,14 +919,23 @@ func (us SqlUserStore) GetByAuth(authData *string, authService string) store.Sto return } - user := model.User{} + query := us.usersQuery. + Where("u.AuthData = ?", authData). + Where("u.AuthService = ?", authService) - if err := us.GetReplica().SelectOne(&user, "SELECT * FROM Users WHERE AuthData = :AuthData AND AuthService = :AuthService", map[string]interface{}{"AuthData": authData, "AuthService": authService}); err != nil { - if err == sql.ErrNoRows { - result.Err = model.NewAppError("SqlUserStore.GetByAuth", store.MISSING_AUTH_ACCOUNT_ERROR, nil, "authData="+*authData+", authService="+authService+", "+err.Error(), http.StatusInternalServerError) - } else { - result.Err = model.NewAppError("SqlUserStore.GetByAuth", "store.sql_user.get_by_auth.other.app_error", nil, "authData="+*authData+", authService="+authService+", "+err.Error(), http.StatusInternalServerError) - } + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetByAuth", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + user := model.User{} + if err := us.GetReplica().SelectOne(&user, queryString, args...); err == sql.ErrNoRows { + result.Err = model.NewAppError("SqlUserStore.GetByAuth", store.MISSING_AUTH_ACCOUNT_ERROR, nil, "authData="+*authData+", authService="+authService+", "+err.Error(), http.StatusInternalServerError) + return + } else if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetByAuth", "store.sql_user.get_by_auth.other.app_error", nil, "authData="+*authData+", authService="+authService+", "+err.Error(), http.StatusInternalServerError) + return } result.Data = &user @@ -878,10 +944,20 @@ func (us SqlUserStore) GetByAuth(authData *string, authService string) store.Sto func (us SqlUserStore) GetAllUsingAuthService(authService string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var data []*model.User + query := us.usersQuery. + Where("u.AuthService = ?", authService). + OrderBy("u.Username ASC") - if _, err := us.GetReplica().Select(&data, "SELECT * FROM Users WHERE AuthService = :AuthService", map[string]interface{}{"AuthService": authService}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetByAuth", "store.sql_user.get_by_auth.other.app_error", nil, "authService="+authService+", "+err.Error(), http.StatusInternalServerError) + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAllUsingAuthService", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + var data []*model.User + if _, err := us.GetReplica().Select(&data, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetAllUsingAuthService", "store.sql_user.get_by_auth.other.app_error", nil, "authService="+authService+", "+err.Error(), http.StatusInternalServerError) + return } result.Data = data @@ -890,10 +966,18 @@ func (us SqlUserStore) GetAllUsingAuthService(authService string) store.StoreCha func (us SqlUserStore) GetByUsername(username string) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - user := model.User{} + query := us.usersQuery.Where("u.Username = ?", username) - if err := us.GetReplica().SelectOne(&user, "SELECT * FROM Users WHERE Username = :Username", map[string]interface{}{"Username": username}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetByUsername", "store.sql_user.get_by_username.app_error", nil, err.Error(), http.StatusInternalServerError) + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetByUsername", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + user := model.User{} + if err := us.GetReplica().SelectOne(&user, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetByUsername", "store.sql_user.get_by_username.app_error", nil, err.Error()+" -- "+queryString, http.StatusInternalServerError) + return } result.Data = &user @@ -902,31 +986,42 @@ func (us SqlUserStore) GetByUsername(username string) store.StoreChannel { func (us SqlUserStore) GetForLogin(loginId string, allowSignInWithUsername, allowSignInWithEmail bool) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - params := map[string]interface{}{ - "LoginId": loginId, - "AllowSignInWithUsername": allowSignInWithUsername, - "AllowSignInWithEmail": allowSignInWithEmail, + query := us.usersQuery + + if allowSignInWithUsername && allowSignInWithEmail { + query = query.Where("Username = ? OR Email = ?", loginId, loginId) + } else if allowSignInWithUsername { + query = query.Where("Username = ?", loginId) + } else if allowSignInWithEmail { + query = query.Where("Email = ?", loginId) + } else { + result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.get_for_login.app_error", nil, "", http.StatusInternalServerError) + return + } + + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } users := []*model.User{} - if _, err := us.GetReplica().Select( - &users, - `SELECT - * - FROM - Users - WHERE - (:AllowSignInWithUsername AND Username = :LoginId) - OR (:AllowSignInWithEmail AND Email = :LoginId)`, - params); err != nil { + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.get_for_login.app_error", nil, err.Error(), http.StatusInternalServerError) - } else if len(users) == 1 { - result.Data = users[0] - } else if len(users) > 1 { - result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.get_for_login.multiple_users", nil, "", http.StatusInternalServerError) - } else { - result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.get_for_login.app_error", nil, "", http.StatusInternalServerError) + return } + + if len(users) == 0 { + result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.get_for_login.app_error", nil, "", http.StatusInternalServerError) + return + } + + if len(users) > 1 { + result.Err = model.NewAppError("SqlUserStore.GetForLogin", "store.sql_user.get_for_login.multiple_users", nil, "", http.StatusInternalServerError) + return + } + + result.Data = users[0] }) } @@ -1029,162 +1124,73 @@ func (us SqlUserStore) GetAnyUnreadPostCountForChannel(userId string, channelId func (us SqlUserStore) Search(teamId string, term string, options *model.UserSearchOptions) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - searchQuery := "" + query := us.usersQuery. + OrderBy("Username ASC"). + Limit(uint64(options.Limit)) - if teamId == "" { - // Id != '' is added because both SEARCH_CLAUSE and INACTIVE_CLAUSE start with an AND - searchQuery = ` - SELECT - * - FROM - Users - WHERE - Id != '' - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Username ASC - LIMIT :Limit` - } else { - searchQuery = ` - SELECT - Users.* - FROM - Users, TeamMembers - WHERE - TeamMembers.TeamId = :TeamId - AND Users.Id = TeamMembers.UserId - AND TeamMembers.DeleteAt = 0 - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Users.Username ASC - LIMIT :Limit` + if teamId != "" { + query = query.Join("TeamMembers tm ON ( tm.UserId = u.Id AND tm.DeleteAt = 0 AND tm.TeamId = ? )", teamId) } - *result = us.performSearch(searchQuery, term, options, map[string]interface{}{ - "TeamId": teamId, - "Limit": options.Limit, - }) - + *result = us.performSearch(query, term, options) }) } func (us SqlUserStore) SearchWithoutTeam(term string, options *model.UserSearchOptions) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - searchQuery := ` - SELECT - * - FROM - Users - WHERE - (SELECT - COUNT(0) - FROM - TeamMembers - WHERE - TeamMembers.UserId = Users.Id - AND TeamMembers.DeleteAt = 0) = 0 - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Username ASC - LIMIT :Limit` - - *result = us.performSearch(searchQuery, term, options, map[string]interface{}{ - "Limit": options.Limit, - }) + query := us.usersQuery. + Where(`( + SELECT + COUNT(0) + FROM + TeamMembers + WHERE + TeamMembers.UserId = u.Id + AND TeamMembers.DeleteAt = 0 + ) = 0`). + OrderBy("u.Username ASC"). + Limit(uint64(options.Limit)) + *result = us.performSearch(query, term, options) }) } func (us SqlUserStore) SearchNotInTeam(notInTeamId string, term string, options *model.UserSearchOptions) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - searchQuery := ` - SELECT - Users.* - FROM Users - LEFT JOIN TeamMembers tm - ON tm.UserId = Users.Id - AND tm.TeamId = :NotInTeamId - WHERE - (tm.UserId IS NULL OR tm.DeleteAt != 0) - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Users.Username ASC - LIMIT :Limit` - - *result = us.performSearch(searchQuery, term, options, map[string]interface{}{ - "NotInTeamId": notInTeamId, - "Limit": options.Limit, - }) + query := us.usersQuery. + LeftJoin("TeamMembers tm ON ( tm.UserId = u.Id AND tm.DeleteAt = 0 AND tm.TeamId = ? )", notInTeamId). + Where("tm.UserId IS NULL"). + OrderBy("u.Username ASC"). + Limit(uint64(options.Limit)) + *result = us.performSearch(query, term, options) }) } func (us SqlUserStore) SearchNotInChannel(teamId string, channelId string, term string, options *model.UserSearchOptions) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - searchQuery := "" - if teamId == "" { - searchQuery = ` - SELECT - Users.* - FROM Users - LEFT JOIN ChannelMembers cm - ON cm.UserId = Users.Id - AND cm.ChannelId = :ChannelId - WHERE - cm.UserId IS NULL - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Users.Username ASC - LIMIT :Limit` - } else { - searchQuery = ` - SELECT - Users.* - FROM Users - INNER JOIN TeamMembers tm - ON tm.UserId = Users.Id - AND tm.TeamId = :TeamId - AND tm.DeleteAt = 0 - LEFT JOIN ChannelMembers cm - ON cm.UserId = Users.Id - AND cm.ChannelId = :ChannelId - WHERE - cm.UserId IS NULL - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Users.Username ASC - LIMIT :Limit` + query := us.usersQuery. + LeftJoin("ChannelMembers cm ON ( cm.UserId = u.Id AND cm.ChannelId = ? )", channelId). + Where("cm.UserId IS NULL"). + OrderBy("Username ASC"). + Limit(uint64(options.Limit)) + + if teamId != "" { + query = query.Join("TeamMembers tm ON ( tm.UserId = u.Id AND tm.DeleteAt = 0 AND tm.TeamId = ? )", teamId) } - *result = us.performSearch(searchQuery, term, options, map[string]interface{}{ - "TeamId": teamId, - "ChannelId": channelId, - "Limit": options.Limit, - }) + *result = us.performSearch(query, term, options) }) } func (us SqlUserStore) SearchInChannel(channelId string, term string, options *model.UserSearchOptions) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - searchQuery := ` - SELECT - Users.* - FROM - Users, ChannelMembers - WHERE - ChannelMembers.ChannelId = :ChannelId - AND ChannelMembers.UserId = Users.Id - SEARCH_CLAUSE - INACTIVE_CLAUSE - ORDER BY Users.Username ASC - LIMIT :Limit - ` - - *result = us.performSearch(searchQuery, term, options, map[string]interface{}{ - "ChannelId": channelId, - "Limit": options.Limit, - }) + query := us.usersQuery. + Join("ChannelMembers cm ON ( cm.UserId = u.Id AND cm.ChannelId = ? )", channelId). + OrderBy("Username ASC"). + Limit(uint64(options.Limit)) + *result = us.performSearch(query, term, options) }) } @@ -1212,31 +1218,25 @@ var spaceFulltextSearchChar = []string{ "@", } -func generateSearchQuery(searchQuery string, terms []string, fields []string, parameters map[string]interface{}, isPostgreSQL bool, role string) string { - searchTerms := []string{} - for i, term := range terms { +func generateSearchQuery(query sq.SelectBuilder, terms []string, fields []string, isPostgreSQL bool) sq.SelectBuilder { + for _, term := range terms { searchFields := []string{} + termArgs := []interface{}{} for _, field := range fields { if isPostgreSQL { - searchFields = append(searchFields, fmt.Sprintf("lower(%s) LIKE lower(%s) escape '*' ", field, fmt.Sprintf(":Term%d", i))) + searchFields = append(searchFields, fmt.Sprintf("lower(%s) LIKE lower(?) escape '*' ", field)) } else { - searchFields = append(searchFields, fmt.Sprintf("%s LIKE %s escape '*' ", field, fmt.Sprintf(":Term%d", i))) + searchFields = append(searchFields, fmt.Sprintf("%s LIKE ? escape '*' ", field)) } + termArgs = append(termArgs, fmt.Sprintf("%s%%", strings.TrimLeft(term, "@"))) } - searchTerms = append(searchTerms, fmt.Sprintf("(%s)", strings.Join(searchFields, " OR "))) - parameters[fmt.Sprintf("Term%d", i)] = fmt.Sprintf("%s%%", strings.TrimLeft(term, "@")) + query = query.Where(fmt.Sprintf("(%s)", strings.Join(searchFields, " OR ")), termArgs...) } - if role != "" { - searchTerms = append(searchTerms, getRoleFilter(isPostgreSQL)) - parameters["Role"] = fmt.Sprintf("%%%s%%", role) - } - - searchClause := strings.Join(searchTerms, " AND ") - return strings.Replace(searchQuery, "SEARCH_CLAUSE", fmt.Sprintf(" AND %s ", searchClause), 1) + return query } -func (us SqlUserStore) performSearch(searchQuery string, term string, options *model.UserSearchOptions, parameters map[string]interface{}) store.StoreResult { +func (us SqlUserStore) performSearch(query sq.SelectBuilder, term string, options *model.UserSearchOptions) store.StoreResult { result := store.StoreResult{} // These chars must be removed from the like query. @@ -1264,27 +1264,26 @@ func (us SqlUserStore) performSearch(searchQuery string, term string, options *m } } - role := "" - if options.Role != "" { - role = options.Role + isPostgreSQL := us.DriverName() == model.DATABASE_DRIVER_POSTGRES + + query = applyRoleFilter(query, options.Role, isPostgreSQL) + + if !options.AllowInactive { + query = query.Where("u.DeleteAt = 0") } - if ok := options.AllowInactive; ok { - searchQuery = strings.Replace(searchQuery, "INACTIVE_CLAUSE", "", 1) - } else { - searchQuery = strings.Replace(searchQuery, "INACTIVE_CLAUSE", "AND Users.DeleteAt = 0", 1) + if strings.TrimSpace(term) != "" { + query = generateSearchQuery(query, strings.Fields(term), searchType, isPostgreSQL) } - if strings.TrimSpace(term) == "" { - searchQuery = strings.Replace(searchQuery, "SEARCH_CLAUSE", "", 1) - } else { - isPostgreSQL := us.DriverName() == model.DATABASE_DRIVER_POSTGRES - searchQuery = generateSearchQuery(searchQuery, strings.Fields(term), searchType, parameters, isPostgreSQL, role) + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.Search", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return result } var users []*model.User - - if _, err := us.GetReplica().Select(&users, searchQuery, parameters); err != nil { + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { result.Err = model.NewAppError("SqlUserStore.Search", "store.sql_user.search.app_error", nil, fmt.Sprintf("term=%v, search_type=%v, %v", term, searchType, err.Error()), http.StatusInternalServerError) } else { @@ -1320,29 +1319,29 @@ func (us SqlUserStore) AnalyticsGetSystemAdminCount() store.StoreChannel { func (us SqlUserStore) GetProfilesNotInTeam(teamId string, offset int, limit int) store.StoreChannel { return store.Do(func(result *store.StoreResult) { - var users []*model.User + query := us.usersQuery. + LeftJoin("TeamMembers tm ON ( tm.UserId = u.Id AND tm.DeleteAt = 0 AND tm.TeamId = ? )", teamId). + Where("tm.UserId IS NULL"). + OrderBy("u.Username ASC"). + Offset(uint64(offset)).Limit(uint64(limit)) - if _, err := us.GetReplica().Select(&users, ` - SELECT - u.* - FROM Users u - LEFT JOIN TeamMembers tm - ON tm.UserId = u.Id - AND tm.TeamId = :TeamId - AND tm.DeleteAt = 0 - WHERE tm.UserId IS NULL - ORDER BY u.Username ASC - LIMIT :Limit OFFSET :Offset - `, map[string]interface{}{"TeamId": teamId, "Offset": offset, "Limit": limit}); err != nil { - result.Err = model.NewAppError("SqlUserStore.GetProfilesNotInTeam", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) - } else { - - for _, u := range users { - u.Sanitize(map[string]bool{}) - } - - result.Data = users + queryString, args, err := query.ToSql() + if err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesNotInTeam", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + return } + + var users []*model.User + if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil { + result.Err = model.NewAppError("SqlUserStore.GetProfilesNotInTeam", "store.sql_user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError) + return + } + + for _, u := range users { + u.Sanitize(map[string]bool{}) + } + + result.Data = users }) } diff --git a/store/storetest/settings.go b/store/storetest/settings.go index c81d61b48f..dad4d58653 100644 --- a/store/storetest/settings.go +++ b/store/storetest/settings.go @@ -37,10 +37,10 @@ func getEnv(name, defaultValue string) string { func log(message string) { verbose := false if verboseFlag := flag.Lookup("test.v"); verboseFlag != nil { - verbose = verboseFlag.Value.String() == "true" + verbose = verboseFlag.Value.String() != "" } if verboseFlag := flag.Lookup("v"); verboseFlag != nil { - verbose = verboseFlag.Value.String() == "true" + verbose = verboseFlag.Value.String() != "" } if verbose { @@ -207,7 +207,7 @@ func MakeSqlSettings(driver string) *model.SqlSettings { panic("unsupported driver " + driver) } - log("Created temporary database " + dbName) + log("Created temporary " + driver + " database " + dbName) return settings } diff --git a/store/storetest/user_store.go b/store/storetest/user_store.go index ce2421cf51..9910aba292 100644 --- a/store/storetest/user_store.go +++ b/store/storetest/user_store.go @@ -70,9 +70,10 @@ func testUserStoreSave(t *testing.T, ss store.Store) { teamId := model.NewId() maxUsersPerTeam := 50 - u1 := model.User{} - u1.Email = MakeEmail() - u1.Username = model.NewId() + u1 := model.User{ + Email: MakeEmail(), + Username: model.NewId(), + } if err := (<-ss.User().Save(&u1)).Err; err != nil { t.Fatal("couldn't save user", err) @@ -85,41 +86,45 @@ func testUserStoreSave(t *testing.T, ss store.Store) { t.Fatal("shouldn't be able to update user from save") } - u1.Id = "" - if err := (<-ss.User().Save(&u1)).Err; err == nil { + u2 := model.User{ + Email: u1.Email, + Username: model.NewId(), + } + if err := (<-ss.User().Save(&u2)).Err; err == nil { t.Fatal("should be unique email") } - u1.Email = "" + u2.Email = MakeEmail() + u2.Username = u1.Username if err := (<-ss.User().Save(&u1)).Err; err == nil { t.Fatal("should be unique username") } - u1.Email = strings.Repeat("0123456789", 20) - u1.Username = "" + u2.Username = "" if err := (<-ss.User().Save(&u1)).Err; err == nil { t.Fatal("should be unique username") } for i := 0; i < 49; i++ { - u1.Id = "" - u1.Email = MakeEmail() - u1.Username = model.NewId() - if err := (<-ss.User().Save(&u1)).Err; err != nil { + u := model.User{ + Email: MakeEmail(), + Username: model.NewId(), + } + if err := (<-ss.User().Save(&u)).Err; err != nil { t.Fatal("couldn't save item", err) } - defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + defer func() { store.Must(ss.User().PermanentDelete(u.Id)) }() - store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, maxUsersPerTeam)) + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u.Id}, maxUsersPerTeam)) } - u1.Id = "" - u1.Email = MakeEmail() - u1.Username = model.NewId() - if err := (<-ss.User().Save(&u1)).Err; err != nil { + u2.Id = "" + u2.Email = MakeEmail() + u2.Username = model.NewId() + if err := (<-ss.User().Save(&u2)).Err; err != nil { t.Fatal("couldn't save item", err) } - defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() if err := (<-ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, maxUsersPerTeam)).Err; err == nil { t.Fatal("should be the limit") @@ -127,15 +132,17 @@ func testUserStoreSave(t *testing.T, ss store.Store) { } func testUserStoreUpdate(t *testing.T, ss store.Store) { - u1 := &model.User{} - u1.Email = MakeEmail() + u1 := &model.User{ + Email: MakeEmail(), + } store.Must(ss.User().Save(u1)) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: model.NewId(), UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - u2.AuthService = "ldap" + u2 := &model.User{ + Email: MakeEmail(), + AuthService: "ldap", + } store.Must(ss.User().Save(u2)) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: model.NewId(), UserId: u2.Id}, -1)) @@ -146,14 +153,16 @@ func testUserStoreUpdate(t *testing.T, ss store.Store) { t.Fatal(err) } - u1.Id = "missing" - if err := (<-ss.User().Update(u1, false)).Err; err == nil { + missing := &model.User{} + if err := (<-ss.User().Update(missing, false)).Err; err == nil { t.Fatal("Update should have failed because of missing key") } - u1.Id = model.NewId() - if err := (<-ss.User().Update(u1, false)).Err; err == nil { - t.Fatal("Update should have faile because id change") + newId := &model.User{ + Id: model.NewId(), + } + if err := (<-ss.User().Update(newId, false)).Err; err == nil { + t.Fatal("Update should have failed because id change") } u2.Email = MakeEmail() @@ -161,10 +170,11 @@ func testUserStoreUpdate(t *testing.T, ss store.Store) { t.Fatal("Update should have failed because you can't modify AD/LDAP fields") } - u3 := &model.User{} - u3.Email = MakeEmail() + u3 := &model.User{ + Email: MakeEmail(), + AuthService: "gitlab", + } oldEmail := u3.Email - u3.AuthService = "gitlab" store.Must(ss.User().Save(u3)) defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: model.NewId(), UserId: u3.Id}, -1)) @@ -239,23 +249,39 @@ func testUserStoreUpdateFailedPasswordAttempts(t *testing.T, ss store.Store) { } func testUserStoreGet(t *testing.T, ss store.Store) { - u1 := &model.User{} - u1.Email = MakeEmail() + u1 := &model.User{ + Email: MakeEmail(), + } store.Must(ss.User().Save(u1)) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: model.NewId(), UserId: u1.Id}, -1)) - if r1 := <-ss.User().Get(u1.Id); r1.Err != nil { - t.Fatal(r1.Err) - } else { - if r1.Data.(*model.User).ToJson() != u1.ToJson() { - t.Fatal("invalid returned user") - } - } + t.Run("fetch empty id", func(t *testing.T) { + require.NotNil(t, (<-ss.User().Get("")).Err) + }) - if err := (<-ss.User().Get("")).Err; err == nil { - t.Fatal("Missing id should have failed") - } + t.Run("fetch user 1", func(t *testing.T) { + result := <-ss.User().Get(u1.Id) + require.Nil(t, result.Err) + + actual := result.Data.(*model.User) + require.Equal(t, u1, actual) + }) + + t.Run("fetch user 2", func(t *testing.T) { + result := <-ss.User().Get(u2.Id) + require.Nil(t, result.Err) + + actual := result.Data.(*model.User) + require.Equal(t, u2, actual) + }) } func testUserCount(t *testing.T, ss store.Store) { @@ -274,545 +300,648 @@ func testUserCount(t *testing.T, ss store.Store) { } func testGetAllUsingAuthService(t *testing.T, ss store.Store) { - u1 := &model.User{} - u1.Email = MakeEmail() - u1.AuthService = "someservice" - store.Must(ss.User().Save(u1)) + teamId := model.NewId() + + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + AuthService: "service", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - u2.AuthService = "someservice" - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + AuthService: "service", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if r1 := <-ss.User().GetAllUsingAuthService(u1.AuthService); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) < 2 { - t.Fatal("invalid returned users") - } - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + AuthService: "service2", + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + + t.Run("get by unknown auth service", func(t *testing.T) { + result := <-ss.User().GetAllUsingAuthService("unknown") + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{}, result.Data.([]*model.User)) + }) + + t.Run("get by auth service", func(t *testing.T) { + result := <-ss.User().GetAllUsingAuthService("service") + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u1, u2}, result.Data.([]*model.User)) + }) + + t.Run("get by other auth service", func(t *testing.T) { + result := <-ss.User().GetAllUsingAuthService("service2") + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u3}, result.Data.([]*model.User)) + }) +} + +func sanitized(user *model.User) *model.User { + clonedUser := model.UserFromJson(strings.NewReader(user.ToJson())) + clonedUser.AuthData = new(string) + *clonedUser.AuthData = "" + clonedUser.Props = model.StringMap{} + + return clonedUser } func testUserStoreGetAllProfiles(t *testing.T, ss store.Store) { - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() - options := &model.UserGetOptions{Page: 0, PerPage: 100} - - if r1 := <-ss.User().GetAllProfiles(options); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) < 2 { - t.Fatal("invalid returned users") - } - } - - options = &model.UserGetOptions{Page: 0, PerPage: 1} - if r2 := <-ss.User().GetAllProfiles(options); r2.Err != nil { - t.Fatal(r2.Err) - } else { - users := r2.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users, limit did not work") - } - } - - if r2 := <-ss.User().GetAll(); r2.Err != nil { - t.Fatal(r2.Err) - } else { - users := r2.Data.([]*model.User) - if len(users) < 2 { - t.Fatal("invalid returned users") - } - } - - etag := "" - if r2 := <-ss.User().GetEtagForAllProfiles(); r2.Err != nil { - t.Fatal(r2.Err) - } else { - etag = r2.Data.(string) - } - - u3 := &model.User{} - u3.Email = MakeEmail() - u3.Roles = "system_user some-other-role" - store.Must(ss.User().Save(u3)) + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() - if r2 := <-ss.User().GetEtagForAllProfiles(); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if etag == r2.Data.(string) { - t.Fatal("etags should not match") - } - } - - u4 := &model.User{} - u4.Email = MakeEmail() - u4.Roles = "system_admin some-other-role" - store.Must(ss.User().Save(u4)) + u4 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u4" + model.NewId(), + Roles: "system_user some-other-role", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u4.Id)) }() - u5 := &model.User{} - u5.Email = MakeEmail() - u5.Roles = "system_admin" - store.Must(ss.User().Save(u5)) + u5 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u5" + model.NewId(), + Roles: "system_admin", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u5.Id)) }() - options = &model.UserGetOptions{Page: 0, PerPage: 10, Role: "system_admin"} - if r2 := <-ss.User().GetAllProfiles(options); r2.Err != nil { - t.Fatal(r2.Err) - } else { - users := r2.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users, role filter did not work") - } - assert.ElementsMatch(t, []string{u4.Id, u5.Id}, []string{users[0].Id, users[1].Id}) - } - - u6 := &model.User{} - u6.Email = MakeEmail() - u6.DeleteAt = model.GetMillis() - u6.Roles = "system_admin" - store.Must(ss.User().Save(u6)) + u6 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u6" + model.NewId(), + DeleteAt: model.GetMillis(), + Roles: "system_admin", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u6.Id)) }() - u7 := &model.User{} - u7.Email = MakeEmail() - u7.DeleteAt = model.GetMillis() - store.Must(ss.User().Save(u7)) + u7 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u7" + model.NewId(), + DeleteAt: model.GetMillis(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u7.Id)) }() - options = &model.UserGetOptions{Page: 0, PerPage: 10, Role: "system_admin", Inactive: true} - if r2 := <-ss.User().GetAllProfiles(options); r2.Err != nil { - t.Fatal(r2.Err) - } else { - users := r2.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users, Role and Inactive filter did not work") - } - assert.Equal(t, u6.Id, users[0].Id) - } + t.Run("get offset 0, limit 100", func(t *testing.T) { + options := &model.UserGetOptions{Page: 0, PerPage: 100} + result := <-ss.User().GetAllProfiles(options) + require.Nil(t, result.Err) - options = &model.UserGetOptions{Page: 0, PerPage: 10, Inactive: true} - if r2 := <-ss.User().GetAllProfiles(options); r2.Err != nil { - t.Fatal(r2.Err) - } else { - users := r2.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users, Inactive filter did not work") - } - assert.ElementsMatch(t, []string{u6.Id, u7.Id}, []string{users[0].Id, users[1].Id}) - } + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u1), + sanitized(u2), + sanitized(u3), + sanitized(u4), + sanitized(u5), + sanitized(u6), + sanitized(u7), + }, actual) + }) + + t.Run("get offset 0, limit 1", func(t *testing.T) { + result := <-ss.User().GetAllProfiles(&model.UserGetOptions{ + Page: 0, + PerPage: 1, + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u1), + }, actual) + }) + + t.Run("get all", func(t *testing.T) { + result := <-ss.User().GetAll() + require.Nil(t, result.Err) + + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + u1, + u2, + u3, + u4, + u5, + u6, + u7, + }, actual) + }) + + t.Run("etag changes for all after user creation", func(t *testing.T) { + result := <-ss.User().GetEtagForAllProfiles() + require.Nil(t, result.Err) + etag := result.Data.(string) + + uNew := &model.User{} + uNew.Email = MakeEmail() + store.Must(ss.User().Save(uNew)) + defer func() { store.Must(ss.User().PermanentDelete(uNew.Id)) }() + + result = <-ss.User().GetEtagForAllProfiles() + require.Nil(t, result.Err) + updatedEtag := result.Data.(string) + + require.NotEqual(t, etag, updatedEtag) + }) + + t.Run("filter to system_admin role", func(t *testing.T) { + result := <-ss.User().GetAllProfiles(&model.UserGetOptions{ + Page: 0, + PerPage: 10, + Role: "system_admin", + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u5), + sanitized(u6), + }, actual) + }) + + t.Run("filter to system_admin role, inactive", func(t *testing.T) { + result := <-ss.User().GetAllProfiles(&model.UserGetOptions{ + Page: 0, + PerPage: 10, + Role: "system_admin", + Inactive: true, + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u6), + }, actual) + }) + + t.Run("filter to inactive", func(t *testing.T) { + result := <-ss.User().GetAllProfiles(&model.UserGetOptions{ + Page: 0, + PerPage: 10, + Inactive: true, + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u6), + sanitized(u7), + }, actual) + }) } func testUserStoreGetProfiles(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - options := &model.UserGetOptions{InTeamId: teamId, Page: 0, PerPage: 100} - if r1 := <-ss.User().GetProfiles(options); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } - - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } - - if !found { - t.Fatal("missing user") - } - } - - options = &model.UserGetOptions{InTeamId: "123", Page: 0, PerPage: 100} - if r2 := <-ss.User().GetProfiles(options); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.([]*model.User)) != 0 { - t.Fatal("should have returned empty map") - } - } - - etag := "" - if r2 := <-ss.User().GetEtagForProfiles(teamId); r2.Err != nil { - t.Fatal(r2.Err) - } else { - etag = r2.Data.(string) - } - - u3 := &model.User{} - u3.Email = MakeEmail() - store.Must(ss.User().Save(u3)) + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - if r2 := <-ss.User().GetEtagForProfiles(teamId); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if etag == r2.Data.(string) { - t.Fatal("etags should not match") - } - } - - u4 := &model.User{} - u4.Email = MakeEmail() - u4.Roles = "system_admin" - store.Must(ss.User().Save(u4)) + u4 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u4" + model.NewId(), + Roles: "system_admin", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u4.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u4.Id}, -1)) - u5 := &model.User{} - u5.Email = MakeEmail() - u5.DeleteAt = model.GetMillis() - store.Must(ss.User().Save(u5)) + u5 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u5" + model.NewId(), + DeleteAt: model.GetMillis(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u5.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u5.Id}, -1)) - options = &model.UserGetOptions{InTeamId: teamId, Page: 0, PerPage: 100} - if r1 := <-ss.User().GetProfiles(options); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 5 { - t.Fatal("invalid returned users") - } - } + t.Run("get page 0, perPage 100", func(t *testing.T) { + result := <-ss.User().GetProfiles(&model.UserGetOptions{ + InTeamId: teamId, + Page: 0, + PerPage: 100, + }) + require.Nil(t, result.Err) - options = &model.UserGetOptions{InTeamId: teamId, Role: "system_admin", Inactive: false, Page: 0, PerPage: 100} - if r1 := <-ss.User().GetProfiles(options); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } - assert.Equal(t, u4.Id, users[0].Id) - } + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u1), + sanitized(u2), + sanitized(u3), + sanitized(u4), + sanitized(u5), + }, actual) + }) - options = &model.UserGetOptions{InTeamId: teamId, Inactive: true, Page: 0, PerPage: 100} - if r1 := <-ss.User().GetProfiles(options); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } - assert.Equal(t, u5.Id, users[0].Id) - } + t.Run("get page 0, perPage 1", func(t *testing.T) { + result := <-ss.User().GetProfiles(&model.UserGetOptions{ + InTeamId: teamId, + Page: 0, + PerPage: 1, + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{sanitized(u1)}, actual) + }) + + t.Run("get unknown team id", func(t *testing.T) { + result := <-ss.User().GetProfiles(&model.UserGetOptions{ + InTeamId: "123", + Page: 0, + PerPage: 100, + }) + require.Nil(t, result.Err) + + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{}, actual) + }) + + t.Run("etag changes for all after user creation", func(t *testing.T) { + result := <-ss.User().GetEtagForProfiles(teamId) + require.Nil(t, result.Err) + etag := result.Data.(string) + + uNew := &model.User{} + uNew.Email = MakeEmail() + store.Must(ss.User().Save(uNew)) + defer func() { store.Must(ss.User().PermanentDelete(uNew.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: uNew.Id}, -1)) + + result = <-ss.User().GetEtagForProfiles(teamId) + require.Nil(t, result.Err) + updatedEtag := result.Data.(string) + + require.NotEqual(t, etag, updatedEtag) + }) + + t.Run("filter to system_admin role", func(t *testing.T) { + result := <-ss.User().GetProfiles(&model.UserGetOptions{ + InTeamId: teamId, + Page: 0, + PerPage: 10, + Role: "system_admin", + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u4), + }, actual) + }) + + t.Run("filter to inactive", func(t *testing.T) { + result := <-ss.User().GetProfiles(&model.UserGetOptions{ + InTeamId: teamId, + Page: 0, + PerPage: 10, + Inactive: true, + }) + require.Nil(t, result.Err) + actual := result.Data.([]*model.User) + require.Equal(t, []*model.User{ + sanitized(u5), + }, actual) + }) } func testUserStoreGetProfilesInChannel(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - c1 := model.Channel{} - c1.TeamId = teamId - c1.DisplayName = "Profiles in channel" - c1.Name = "profiles-" + model.NewId() - c1.Type = model.CHANNEL_OPEN + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - c2 := model.Channel{} - c2.TeamId = teamId - c2.DisplayName = "Profiles in private" - c2.Name = "profiles-" + model.NewId() - c2.Type = model.CHANNEL_PRIVATE + c1 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in channel", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_OPEN, + }, -1)).(*model.Channel) - store.Must(ss.Channel().Save(&c1, -1)) - store.Must(ss.Channel().Save(&c2, -1)) + c2 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in private", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_PRIVATE, + }, -1)).(*model.Channel) - m1 := model.ChannelMember{} - m1.ChannelId = c1.Id - m1.UserId = u1.Id - m1.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m2 := model.ChannelMember{} - m2.ChannelId = c1.Id - m2.UserId = u2.Id - m2.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u2.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m3 := model.ChannelMember{} - m3.ChannelId = c2.Id - m3.UserId = u1.Id - m3.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u3.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - store.Must(ss.Channel().SaveMember(&m1)) - store.Must(ss.Channel().SaveMember(&m2)) - store.Must(ss.Channel().SaveMember(&m3)) + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c2.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - if r1 := <-ss.User().GetProfilesInChannel(c1.Id, 0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } + t.Run("get in channel 1, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetProfilesInChannel(c1.Id, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1), sanitized(u2), sanitized(u3)}, result.Data.([]*model.User)) + }) - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } + t.Run("get in channel 1, offset 1, limit 2", func(t *testing.T) { + result := <-ss.User().GetProfilesInChannel(c1.Id, 1, 2) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u2), sanitized(u3)}, result.Data.([]*model.User)) + }) - if !found { - t.Fatal("missing user") - } - } - - if r2 := <-ss.User().GetProfilesInChannel(c2.Id, 0, 1); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.([]*model.User)) != 1 { - t.Fatal("should have returned only 1 user") - } - } + t.Run("get in channel 2, offset 0, limit 1", func(t *testing.T) { + result := <-ss.User().GetProfilesInChannel(c2.Id, 0, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1)}, result.Data.([]*model.User)) + }) } func testUserStoreGetProfilesInChannelByStatus(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - c1 := model.Channel{} - c1.TeamId = teamId - c1.DisplayName = "Profiles in channel" - c1.Name = "profiles-" + model.NewId() - c1.Type = model.CHANNEL_OPEN + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - c2 := model.Channel{} - c2.TeamId = teamId - c2.DisplayName = "Profiles in private" - c2.Name = "profiles-" + model.NewId() - c2.Type = model.CHANNEL_PRIVATE + c1 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in channel", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_OPEN, + }, -1)).(*model.Channel) - store.Must(ss.Channel().Save(&c1, -1)) - store.Must(ss.Channel().Save(&c2, -1)) + c2 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in private", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_PRIVATE, + }, -1)).(*model.Channel) - m1 := model.ChannelMember{} - m1.ChannelId = c1.Id - m1.UserId = u1.Id - m1.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m2 := model.ChannelMember{} - m2.ChannelId = c1.Id - m2.UserId = u2.Id - m2.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u2.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m3 := model.ChannelMember{} - m3.ChannelId = c2.Id - m3.UserId = u1.Id - m3.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u3.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - store.Must(ss.Channel().SaveMember(&m1)) - store.Must(ss.Channel().SaveMember(&m2)) - store.Must(ss.Channel().SaveMember(&m3)) + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c2.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - if r1 := <-ss.User().GetProfilesInChannelByStatus(c1.Id, 0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } + store.Must(ss.Status().SaveOrUpdate(&model.Status{ + UserId: u1.Id, + Status: model.STATUS_DND, + })) + store.Must(ss.Status().SaveOrUpdate(&model.Status{ + UserId: u2.Id, + Status: model.STATUS_AWAY, + })) + store.Must(ss.Status().SaveOrUpdate(&model.Status{ + UserId: u3.Id, + Status: model.STATUS_ONLINE, + })) - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } + t.Run("get in channel 1 by status, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetProfilesInChannelByStatus(c1.Id, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u3), sanitized(u2), sanitized(u1)}, result.Data.([]*model.User)) + }) - if !found { - t.Fatal("missing user") - } - } - - if r2 := <-ss.User().GetProfilesInChannelByStatus(c2.Id, 0, 1); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.([]*model.User)) != 1 { - t.Fatal("should have returned only 1 user") - } - } + t.Run("get in channel 2 by status, offset 0, limit 1", func(t *testing.T) { + result := <-ss.User().GetProfilesInChannelByStatus(c2.Id, 0, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1)}, result.Data.([]*model.User)) + }) } func testUserStoreGetProfilesWithoutTeam(t *testing.T, ss store.Store) { teamId := model.NewId() - // These usernames need to appear in the first 100 users for this to work - - u1 := &model.User{} - u1.Username = "a000000000" + model.NewId() - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) - store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Username = "a000000001" + model.NewId() - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() - if r1 := <-ss.User().GetProfilesWithoutTeam(0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() - found1 := false - found2 := false - for _, u := range users { - if u.Id == u1.Id { - found1 = true - } else if u.Id == u2.Id { - found2 = true - } - } + t.Run("get, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetProfilesWithoutTeam(0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u2), sanitized(u3)}, result.Data.([]*model.User)) + }) - if found1 { - t.Fatal("shouldn't have returned user on team") - } else if !found2 { - t.Fatal("should've returned user without any teams") - } - } + t.Run("get, offset 1, limit 1", func(t *testing.T) { + result := <-ss.User().GetProfilesWithoutTeam(1, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u3)}, result.Data.([]*model.User)) + }) + + t.Run("get, offset 2, limit 1", func(t *testing.T) { + result := <-ss.User().GetProfilesWithoutTeam(2, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{}, result.Data.([]*model.User)) + }) } func testUserStoreGetAllProfilesInChannel(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - c1 := model.Channel{} - c1.TeamId = teamId - c1.DisplayName = "Profiles in channel" - c1.Name = "profiles-" + model.NewId() - c1.Type = model.CHANNEL_OPEN + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - c2 := model.Channel{} - c2.TeamId = teamId - c2.DisplayName = "Profiles in private" - c2.Name = "profiles-" + model.NewId() - c2.Type = model.CHANNEL_PRIVATE + c1 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in channel", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_OPEN, + }, -1)).(*model.Channel) - store.Must(ss.Channel().Save(&c1, -1)) - store.Must(ss.Channel().Save(&c2, -1)) + c2 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in private", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_PRIVATE, + }, -1)).(*model.Channel) - m1 := model.ChannelMember{} - m1.ChannelId = c1.Id - m1.UserId = u1.Id - m1.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m2 := model.ChannelMember{} - m2.ChannelId = c1.Id - m2.UserId = u2.Id - m2.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u2.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m3 := model.ChannelMember{} - m3.ChannelId = c2.Id - m3.UserId = u1.Id - m3.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u3.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - store.Must(ss.Channel().SaveMember(&m1)) - store.Must(ss.Channel().SaveMember(&m2)) - store.Must(ss.Channel().SaveMember(&m3)) + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c2.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - if r1 := <-ss.User().GetAllProfilesInChannel(c1.Id, false); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.(map[string]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } + t.Run("all profiles in channel 1, no caching", func(t *testing.T) { + result := <-ss.User().GetAllProfilesInChannel(c1.Id, false) + require.Nil(t, result.Err) + assert.Equal(t, map[string]*model.User{ + u1.Id: sanitized(u1), + u2.Id: sanitized(u2), + u3.Id: sanitized(u3), + }, result.Data.(map[string]*model.User)) + }) - if users[u1.Id].Id != u1.Id { - t.Fatal("invalid returned user") - } - } + t.Run("all profiles in channel 2, no caching", func(t *testing.T) { + result := <-ss.User().GetAllProfilesInChannel(c2.Id, false) + require.Nil(t, result.Err) + assert.Equal(t, map[string]*model.User{ + u1.Id: sanitized(u1), + }, result.Data.(map[string]*model.User)) + }) - if r2 := <-ss.User().GetAllProfilesInChannel(c2.Id, false); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.(map[string]*model.User)) != 1 { - t.Fatal("should have returned empty map") - } - } + t.Run("all profiles in channel 2, caching", func(t *testing.T) { + result := <-ss.User().GetAllProfilesInChannel(c2.Id, true) + require.Nil(t, result.Err) + assert.Equal(t, map[string]*model.User{ + u1.Id: sanitized(u1), + }, result.Data.(map[string]*model.User)) + }) - if r2 := <-ss.User().GetAllProfilesInChannel(c2.Id, true); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.(map[string]*model.User)) != 1 { - t.Fatal("should have returned empty map") - } - } - - if r2 := <-ss.User().GetAllProfilesInChannel(c2.Id, true); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.(map[string]*model.User)) != 1 { - t.Fatal("should have returned empty map") - } - } + t.Run("all profiles in channel 2, caching [repeated]", func(t *testing.T) { + result := <-ss.User().GetAllProfilesInChannel(c2.Id, true) + require.Nil(t, result.Err) + assert.Equal(t, map[string]*model.User{ + u1.Id: sanitized(u1), + }, result.Data.(map[string]*model.User)) + }) ss.User().InvalidateProfilesInChannelCacheByUser(u1.Id) ss.User().InvalidateProfilesInChannelCache(c2.Id) @@ -821,459 +950,495 @@ func testUserStoreGetAllProfilesInChannel(t *testing.T, ss store.Store) { func testUserStoreGetProfilesNotInChannel(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - c1 := model.Channel{} - c1.TeamId = teamId - c1.DisplayName = "Profiles in channel" - c1.Name = "profiles-" + model.NewId() - c1.Type = model.CHANNEL_OPEN + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - c2 := model.Channel{} - c2.TeamId = teamId - c2.DisplayName = "Profiles in private" - c2.Name = "profiles-" + model.NewId() - c2.Type = model.CHANNEL_PRIVATE + c1 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in channel", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_OPEN, + }, -1)).(*model.Channel) - store.Must(ss.Channel().Save(&c1, -1)) - store.Must(ss.Channel().Save(&c2, -1)) + c2 := store.Must(ss.Channel().Save(&model.Channel{ + TeamId: teamId, + DisplayName: "Profiles in private", + Name: "profiles-" + model.NewId(), + Type: model.CHANNEL_PRIVATE, + }, -1)).(*model.Channel) - if r1 := <-ss.User().GetProfilesNotInChannel(teamId, c1.Id, 0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } + t.Run("get team 1, channel 1, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInChannel(teamId, c1.Id, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u1), + sanitized(u2), + sanitized(u3), + }, result.Data.([]*model.User)) + }) - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } + t.Run("get team 1, channel 2, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInChannel(teamId, c2.Id, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u1), + sanitized(u2), + sanitized(u3), + }, result.Data.([]*model.User)) + }) - if !found { - t.Fatal("missing user") - } - } + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - if r2 := <-ss.User().GetProfilesNotInChannel(teamId, c2.Id, 0, 100); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.([]*model.User)) != 2 { - t.Fatal("invalid returned users") - } - } + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u2.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m1 := model.ChannelMember{} - m1.ChannelId = c1.Id - m1.UserId = u1.Id - m1.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c1.Id, + UserId: u3.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m2 := model.ChannelMember{} - m2.ChannelId = c1.Id - m2.UserId = u2.Id - m2.NotifyProps = model.GetDefaultChannelNotifyProps() + store.Must(ss.Channel().SaveMember(&model.ChannelMember{ + ChannelId: c2.Id, + UserId: u1.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + })) - m3 := model.ChannelMember{} - m3.ChannelId = c2.Id - m3.UserId = u1.Id - m3.NotifyProps = model.GetDefaultChannelNotifyProps() + t.Run("get team 1, channel 1, offset 0, limit 100, after update", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInChannel(teamId, c1.Id, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{}, result.Data.([]*model.User)) + }) - store.Must(ss.Channel().SaveMember(&m1)) - store.Must(ss.Channel().SaveMember(&m2)) - store.Must(ss.Channel().SaveMember(&m3)) - - if r1 := <-ss.User().GetProfilesNotInChannel(teamId, c1.Id, 0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 0 { - t.Fatal("invalid returned users") - } - } - - if r2 := <-ss.User().GetProfilesNotInChannel(teamId, c2.Id, 0, 100); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.([]*model.User)) != 1 { - t.Fatal("should have had 1 user not in channel") - } - } + t.Run("get team 1, channel 2, offset 0, limit 100, after update", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInChannel(teamId, c2.Id, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u2), + sanitized(u3), + }, result.Data.([]*model.User)) + }) } func testUserStoreGetProfilesByIds(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if r1 := <-ss.User().GetProfileByIds([]string{u1.Id}, false); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } + t.Run("get u1 by id, no caching", func(t *testing.T) { + result := <-ss.User().GetProfileByIds([]string{u1.Id}, false) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1)}, result.Data.([]*model.User)) + }) - if !found { - t.Fatal("missing user") - } - } + t.Run("get u1 by id, caching", func(t *testing.T) { + result := <-ss.User().GetProfileByIds([]string{u1.Id}, true) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1)}, result.Data.([]*model.User)) + }) - if r1 := <-ss.User().GetProfileByIds([]string{u1.Id}, true); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } + t.Run("get u1, u2, u3 by id, no caching", func(t *testing.T) { + result := <-ss.User().GetProfileByIds([]string{u1.Id, u2.Id, u3.Id}, false) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1), sanitized(u2), sanitized(u3)}, result.Data.([]*model.User)) + }) - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } + t.Run("get u1, u2, u3 by id, caching", func(t *testing.T) { + result := <-ss.User().GetProfileByIds([]string{u1.Id, u2.Id, u3.Id}, true) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{sanitized(u1), sanitized(u2), sanitized(u3)}, result.Data.([]*model.User)) + }) - if !found { - t.Fatal("missing user") - } - } - - if r1 := <-ss.User().GetProfileByIds([]string{u1.Id, u2.Id}, true); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } - - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } - - if !found { - t.Fatal("missing user") - } - } - - if r1 := <-ss.User().GetProfileByIds([]string{u1.Id, u2.Id}, true); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } - - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } - - if !found { - t.Fatal("missing user") - } - } - - if r1 := <-ss.User().GetProfileByIds([]string{u1.Id, u2.Id}, false); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } - - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } - - if !found { - t.Fatal("missing user") - } - } - - if r1 := <-ss.User().GetProfileByIds([]string{u1.Id}, false); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } - - found := false - for _, u := range users { - if u.Id == u1.Id { - found = true - } - } - - if !found { - t.Fatal("missing user") - } - } - - options := &model.UserGetOptions{InTeamId: "123", Page: 0, PerPage: 100} - if r2 := <-ss.User().GetProfiles(options); r2.Err != nil { - t.Fatal(r2.Err) - } else { - if len(r2.Data.([]*model.User)) != 0 { - t.Fatal("should have returned empty array") - } - } + t.Run("get unknown id, caching", func(t *testing.T) { + result := <-ss.User().GetProfileByIds([]string{"123"}, true) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{}, result.Data.([]*model.User)) + }) } func testUserStoreGetProfilesByUsernames(t *testing.T, ss store.Store) { teamId := model.NewId() + team2Id := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - u1.Username = "username1" + model.NewId() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - u2.Username = "username2" + model.NewId() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if r1 := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u2.Username}, teamId); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } - - if users[0].Id != u1.Id && users[1].Id != u1.Id { - t.Fatal("invalid returned user 1") - } - - if users[0].Id != u2.Id && users[1].Id != u2.Id { - t.Fatal("invalid returned user 2") - } - } - - if r1 := <-ss.User().GetProfilesByUsernames([]string{u1.Username}, teamId); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } - - if users[0].Id != u1.Id { - t.Fatal("invalid returned user") - } - } - - team2Id := model.NewId() - - u3 := &model.User{} - u3.Email = MakeEmail() - u3.Username = "username3" + model.NewId() - store.Must(ss.User().Save(u3)) + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: team2Id, UserId: u3.Id}, -1)) - if r1 := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u3.Username}, ""); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 2 { - t.Fatal("invalid returned users") - } + t.Run("get by u1 and u2 usernames, team id 1", func(t *testing.T) { + result := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u2.Username}, teamId) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u1, u2}, result.Data.([]*model.User)) + }) - if users[0].Id != u1.Id && users[1].Id != u1.Id { - t.Fatal("invalid returned user 1") - } + t.Run("get by u1 username, team id 1", func(t *testing.T) { + result := <-ss.User().GetProfilesByUsernames([]string{u1.Username}, teamId) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u1}, result.Data.([]*model.User)) + }) - if users[0].Id != u3.Id && users[1].Id != u3.Id { - t.Fatal("invalid returned user 3") - } - } + t.Run("get by u1 and u3 usernames, no team id", func(t *testing.T) { + result := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u3.Username}, "") + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u1, u3}, result.Data.([]*model.User)) + }) - if r1 := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u3.Username}, teamId); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - if len(users) != 1 { - t.Fatal("invalid returned users") - } + t.Run("get by u1 and u3 usernames, team id 1", func(t *testing.T) { + result := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u3.Username}, teamId) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u1}, result.Data.([]*model.User)) + }) - if users[0].Id != u1.Id { - t.Fatal("invalid returned user") - } - } + t.Run("get by u1 and u3 usernames, team id 2", func(t *testing.T) { + result := <-ss.User().GetProfilesByUsernames([]string{u1.Username, u3.Username}, team2Id) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{u3}, result.Data.([]*model.User)) + }) } func testUserStoreGetSystemAdminProfiles(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - u1.Roles = model.SYSTEM_USER_ROLE_ID + " " + model.SYSTEM_ADMIN_ROLE_ID - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Roles: model.SYSTEM_USER_ROLE_ID + " " + model.SYSTEM_ADMIN_ROLE_ID, + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if r1 := <-ss.User().GetSystemAdminProfiles(); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.(map[string]*model.User) - if len(users) <= 0 { - t.Fatal("invalid returned system admin users") - } - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Roles: model.SYSTEM_USER_ROLE_ID + " " + model.SYSTEM_ADMIN_ROLE_ID, + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + + t.Run("all system admin profiles", func(t *testing.T) { + result := <-ss.User().GetSystemAdminProfiles() + require.Nil(t, result.Err) + assert.Equal(t, map[string]*model.User{ + u1.Id: sanitized(u1), + u3.Id: sanitized(u3), + }, result.Data.(map[string]*model.User)) + }) } func testUserStoreGetByEmail(t *testing.T, ss store.Store) { - teamid := model.NewId() + teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() - store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamid, UserId: u1.Id}, -1)) + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - if err := (<-ss.User().GetByEmail(u1.Email)).Err; err != nil { - t.Fatal(err) - } + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if err := (<-ss.User().GetByEmail("")).Err; err == nil { - t.Fatal("Should have failed because of missing email") - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + + t.Run("get u1 by email", func(t *testing.T) { + result := <-ss.User().GetByEmail(u1.Email) + require.Nil(t, result.Err) + assert.Equal(t, u1, result.Data.(*model.User)) + }) + + t.Run("get u2 by email", func(t *testing.T) { + result := <-ss.User().GetByEmail(u2.Email) + require.Nil(t, result.Err) + assert.Equal(t, u2, result.Data.(*model.User)) + }) + + t.Run("get u3 by email", func(t *testing.T) { + result := <-ss.User().GetByEmail(u3.Email) + require.Nil(t, result.Err) + assert.Equal(t, u3, result.Data.(*model.User)) + }) + + t.Run("get by empty email", func(t *testing.T) { + result := <-ss.User().GetByEmail("") + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, store.MISSING_ACCOUNT_ERROR) + }) + + t.Run("get by unknown", func(t *testing.T) { + result := <-ss.User().GetByEmail("unknown") + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, store.MISSING_ACCOUNT_ERROR) + }) } func testUserStoreGetByAuthData(t *testing.T, ss store.Store) { teamId := model.NewId() + auth1 := model.NewId() + auth3 := model.NewId() - auth := "123" + model.NewId() - - u1 := &model.User{} - u1.Email = MakeEmail() - u1.AuthData = &auth - u1.AuthService = "service" - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + AuthData: &auth1, + AuthService: "service", + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - if err := (<-ss.User().GetByAuth(u1.AuthData, u1.AuthService)).Err; err != nil { - t.Fatal(err) - } + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - rauth := "" - if err := (<-ss.User().GetByAuth(&rauth, "")).Err; err == nil { - t.Fatal("Should have failed because of missing auth data") - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + AuthData: &auth3, + AuthService: "service2", + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + + t.Run("get by u1 auth", func(t *testing.T) { + result := <-ss.User().GetByAuth(u1.AuthData, u1.AuthService) + require.Nil(t, result.Err) + assert.Equal(t, u1, result.Data.(*model.User)) + }) + + t.Run("get by u3 auth", func(t *testing.T) { + result := <-ss.User().GetByAuth(u3.AuthData, u3.AuthService) + require.Nil(t, result.Err) + assert.Equal(t, u3, result.Data.(*model.User)) + }) + + t.Run("get by u1 auth, unknown service", func(t *testing.T) { + result := <-ss.User().GetByAuth(u1.AuthData, "unknown") + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, store.MISSING_AUTH_ACCOUNT_ERROR) + }) + + t.Run("get by unknown auth, u1 service", func(t *testing.T) { + unknownAuth := "" + result := <-ss.User().GetByAuth(&unknownAuth, u1.AuthService) + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, store.MISSING_AUTH_ACCOUNT_ERROR) + }) + + t.Run("get by unknown auth, unknown service", func(t *testing.T) { + unknownAuth := "" + result := <-ss.User().GetByAuth(&unknownAuth, "unknown") + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, store.MISSING_AUTH_ACCOUNT_ERROR) + }) } func testUserStoreGetByUsername(t *testing.T, ss store.Store) { teamId := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - u1.Username = model.NewId() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - if err := (<-ss.User().GetByUsername(u1.Username)).Err; err != nil { - t.Fatal(err) - } + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if err := (<-ss.User().GetByUsername("")).Err; err == nil { - t.Fatal("Should have failed because of missing username") - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + + t.Run("get u1 by username", func(t *testing.T) { + result := <-ss.User().GetByUsername(u1.Username) + require.Nil(t, result.Err) + assert.Equal(t, u1, result.Data.(*model.User)) + }) + + t.Run("get u2 by username", func(t *testing.T) { + result := <-ss.User().GetByUsername(u2.Username) + require.Nil(t, result.Err) + assert.Equal(t, u2, result.Data.(*model.User)) + }) + + t.Run("get u3 by username", func(t *testing.T) { + result := <-ss.User().GetByUsername(u3.Username) + require.Nil(t, result.Err) + assert.Equal(t, u3, result.Data.(*model.User)) + }) + + t.Run("get by empty username", func(t *testing.T) { + result := <-ss.User().GetByUsername("") + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, "store.sql_user.get_by_username.app_error") + }) + + t.Run("get by unknown", func(t *testing.T) { + result := <-ss.User().GetByUsername("unknown") + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, "store.sql_user.get_by_username.app_error") + }) } func testUserStoreGetForLogin(t *testing.T, ss store.Store) { + teamId := model.NewId() auth := model.NewId() + auth2 := model.NewId() + auth3 := model.NewId() - u1 := &model.User{ + u1 := store.Must(ss.User().Save(&model.User{ Email: MakeEmail(), - Username: model.NewId(), + Username: "u1" + model.NewId(), AuthService: model.USER_AUTH_SERVICE_GITLAB, AuthData: &auth, - } - store.Must(ss.User().Save(u1)) + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - auth2 := model.NewId() - - u2 := &model.User{ + u2 := store.Must(ss.User().Save(&model.User{ Email: MakeEmail(), - Username: model.NewId(), + Username: "u2" + model.NewId(), AuthService: model.USER_AUTH_SERVICE_LDAP, AuthData: &auth2, - } - store.Must(ss.User().Save(u2)) + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - if result := <-ss.User().GetForLogin(u1.Username, true, true); result.Err != nil { - t.Fatal("Should have gotten user by username", result.Err) - } else if result.Data.(*model.User).Id != u1.Id { - t.Fatal("Should have gotten user1 by username") - } + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + AuthService: model.USER_AUTH_SERVICE_LDAP, + AuthData: &auth3, + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - if result := <-ss.User().GetForLogin(u1.Email, true, true); result.Err != nil { - t.Fatal("Should have gotten user by email", result.Err) - } else if result.Data.(*model.User).Id != u1.Id { - t.Fatal("Should have gotten user1 by email") - } + t.Run("get u1 by username, allow both", func(t *testing.T) { + result := <-ss.User().GetForLogin(u1.Username, true, true) + require.Nil(t, result.Err) + assert.Equal(t, u1, result.Data.(*model.User)) + }) - // prevent getting user when different login methods are disabled - if result := <-ss.User().GetForLogin(u1.Username, false, true); result.Err == nil { - t.Fatal("Should have failed to get user1 by username") - } + t.Run("get u1 by username, allow only email", func(t *testing.T) { + result := <-ss.User().GetForLogin(u1.Username, false, true) + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, "store.sql_user.get_for_login.app_error") + }) - if result := <-ss.User().GetForLogin(u1.Email, true, false); result.Err == nil { - t.Fatal("Should have failed to get user1 by email") - } + t.Run("get u1 by email, allow both", func(t *testing.T) { + result := <-ss.User().GetForLogin(u1.Email, true, true) + require.Nil(t, result.Err) + assert.Equal(t, u1, result.Data.(*model.User)) + }) + + t.Run("get u1 by email, allow only username", func(t *testing.T) { + result := <-ss.User().GetForLogin(u1.Email, true, false) + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, "store.sql_user.get_for_login.app_error") + }) + + t.Run("get u2 by username, allow both", func(t *testing.T) { + result := <-ss.User().GetForLogin(u2.Username, true, true) + require.Nil(t, result.Err) + assert.Equal(t, u2, result.Data.(*model.User)) + }) + + t.Run("get u2 by email, allow both", func(t *testing.T) { + result := <-ss.User().GetForLogin(u2.Email, true, true) + require.Nil(t, result.Err) + assert.Equal(t, u2, result.Data.(*model.User)) + }) + + t.Run("get u2 by username, allow neither", func(t *testing.T) { + result := <-ss.User().GetForLogin(u2.Username, false, false) + require.NotNil(t, result.Err) + require.Equal(t, result.Err.Id, "store.sql_user.get_for_login.app_error") + }) } func testUserStoreUpdatePassword(t *testing.T, ss store.Store) { @@ -1479,31 +1644,130 @@ func testUserStoreUpdateMfaActive(t *testing.T, ss store.Store) { } func testUserStoreGetRecentlyActiveUsersForTeam(t *testing.T, ss store.Store) { - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) - defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() - store.Must(ss.Status().SaveOrUpdate(&model.Status{UserId: u1.Id, Status: model.STATUS_ONLINE, Manual: false, LastActivityAt: model.GetMillis(), ActiveChannel: ""})) - tid := model.NewId() - store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: tid, UserId: u1.Id}, -1)) + teamId := model.NewId() - if r1 := <-ss.User().GetRecentlyActiveUsersForTeam(tid, 0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) + + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) + + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + + millis := model.GetMillis() + u3.LastActivityAt = millis + u2.LastActivityAt = millis - 1 + u1.LastActivityAt = millis - 1 + + store.Must(ss.Status().SaveOrUpdate(&model.Status{UserId: u1.Id, Status: model.STATUS_ONLINE, Manual: false, LastActivityAt: u1.LastActivityAt, ActiveChannel: ""})) + store.Must(ss.Status().SaveOrUpdate(&model.Status{UserId: u2.Id, Status: model.STATUS_ONLINE, Manual: false, LastActivityAt: u2.LastActivityAt, ActiveChannel: ""})) + store.Must(ss.Status().SaveOrUpdate(&model.Status{UserId: u3.Id, Status: model.STATUS_ONLINE, Manual: false, LastActivityAt: u3.LastActivityAt, ActiveChannel: ""})) + + t.Run("get team 1, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetRecentlyActiveUsersForTeam(teamId, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u3), + sanitized(u1), + sanitized(u2), + }, result.Data.([]*model.User)) + }) + + t.Run("get team 1, offset 0, limit 1", func(t *testing.T) { + result := <-ss.User().GetRecentlyActiveUsersForTeam(teamId, 0, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u3), + }, result.Data.([]*model.User)) + }) + + t.Run("get team 1, offset 2, limit 1", func(t *testing.T) { + result := <-ss.User().GetRecentlyActiveUsersForTeam(teamId, 2, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u2), + }, result.Data.([]*model.User)) + }) } func testUserStoreGetNewUsersForTeam(t *testing.T, ss store.Store) { - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) - defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() - store.Must(ss.Status().SaveOrUpdate(&model.Status{UserId: u1.Id, Status: model.STATUS_ONLINE, Manual: false, LastActivityAt: model.GetMillis(), ActiveChannel: ""})) - tid := model.NewId() - store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: tid, UserId: u1.Id}, -1)) + teamId := model.NewId() + teamId2 := model.NewId() - if r1 := <-ss.User().GetNewUsersForTeam(tid, 0, 100); r1.Err != nil { - t.Fatal(r1.Err) - } + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) + + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) + + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) + + u4 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u4" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u4.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId2, UserId: u4.Id}, -1)) + + t.Run("get team 1, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetNewUsersForTeam(teamId, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u3), + sanitized(u2), + sanitized(u1), + }, result.Data.([]*model.User)) + }) + + t.Run("get team 1, offset 0, limit 1", func(t *testing.T) { + result := <-ss.User().GetNewUsersForTeam(teamId, 0, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u3), + }, result.Data.([]*model.User)) + }) + + t.Run("get team 1, offset 2, limit 1", func(t *testing.T) { + result := <-ss.User().GetNewUsersForTeam(teamId, 2, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u1), + }, result.Data.([]*model.User)) + }) + + t.Run("get team 2, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetNewUsersForTeam(teamId2, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u4), + }, result.Data.([]*model.User)) + }) } func assertUsers(t *testing.T, expected, actual []*model.User) { @@ -2558,137 +2822,134 @@ func testUserStoreAnalyticsGetSystemAdminCount(t *testing.T, ss store.Store) { func testUserStoreGetProfilesNotInTeam(t *testing.T, ss store.Store) { teamId := model.NewId() + teamId2 := model.NewId() - u1 := &model.User{} - u1.Email = MakeEmail() - store.Must(ss.User().Save(u1)) + u1 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u1" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u1.Id}, -1)) - store.Must(ss.User().UpdateUpdateAt(u1.Id)) - u2 := &model.User{} - u2.Email = MakeEmail() - store.Must(ss.User().Save(u2)) + // Ensure update at timestamp changes + time.Sleep(time.Millisecond * 10) + + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() - store.Must(ss.User().UpdateUpdateAt(u2.Id)) + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId2, UserId: u2.Id}, -1)) + + // Ensure update at timestamp changes + time.Sleep(time.Millisecond * 10) + + u3 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u3" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() - var initialUsersNotInTeam int var etag1, etag2, etag3 string - if er1 := <-ss.User().GetEtagForProfilesNotInTeam(teamId); er1.Err != nil { - t.Fatal(er1.Err) - } else { - etag1 = er1.Data.(string) - } + t.Run("etag for profiles not in team 1", func(t *testing.T) { + result := <-ss.User().GetEtagForProfilesNotInTeam(teamId) + require.Nil(t, result.Err) + etag1 = result.Data.(string) + }) - if r1 := <-ss.User().GetProfilesNotInTeam(teamId, 0, 100000); r1.Err != nil { - t.Fatal(r1.Err) - } else { - users := r1.Data.([]*model.User) - initialUsersNotInTeam = len(users) - if initialUsersNotInTeam < 1 { - t.Fatalf("Should be at least 1 user not in the team") - } + t.Run("get not in team 1, offset 0, limit 100000", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInTeam(teamId, 0, 100000) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u2), + sanitized(u3), + }, result.Data.([]*model.User)) + }) - found := false - for _, u := range users { - if u.Id == u2.Id { - found = true - } - if u.Id == u1.Id { - t.Fatalf("Should not have found user1") - } - } + t.Run("get not in team 1, offset 1, limit 1", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInTeam(teamId, 1, 1) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u3), + }, result.Data.([]*model.User)) + }) - if !found { - t.Fatal("missing user2") - } - } + t.Run("get not in team 2, offset 0, limit 100", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInTeam(teamId2, 0, 100) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u1), + sanitized(u3), + }, result.Data.([]*model.User)) + }) + // Ensure update at timestamp changes time.Sleep(time.Millisecond * 10) + + // Add u2 to team 1 store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u2.Id}, -1)) - store.Must(ss.User().UpdateUpdateAt(u2.Id)) + u2.UpdateAt = store.Must(ss.User().UpdateUpdateAt(u2.Id)).(int64) - if er2 := <-ss.User().GetEtagForProfilesNotInTeam(teamId); er2.Err != nil { - t.Fatal(er2.Err) - } else { - etag2 = er2.Data.(string) - if etag1 == etag2 { - t.Fatalf("etag should have changed") - } - } + // GetEtagForProfilesNotInTeam only works if the most recent user is added to the team, + // otherwise the timestamp simply never changes: see https://mattermost.atlassian.net/browse/MM-13721. + t.Run("etag for profiles not in team 1 after update", func(t *testing.T) { + t.Skip() + result := <-ss.User().GetEtagForProfilesNotInTeam(teamId) + require.Nil(t, result.Err) + etag2 = result.Data.(string) + require.NotEqual(t, etag2, etag1, "etag should have changed") + }) - if r2 := <-ss.User().GetProfilesNotInTeam(teamId, 0, 100000); r2.Err != nil { - t.Fatal(r2.Err) - } else { - users := r2.Data.([]*model.User) - - if len(users) != initialUsersNotInTeam-1 { - t.Fatalf("Should be one less user not in team") - } - - for _, u := range users { - if u.Id == u2.Id { - t.Fatalf("Should not have found user2") - } - if u.Id == u1.Id { - t.Fatalf("Should not have found user1") - } - } - } + t.Run("get not in team 1, offset 0, limit 100000 after update", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInTeam(teamId, 0, 100000) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u3), + }, result.Data.([]*model.User)) + }) + // Ensure update at timestamp changes time.Sleep(time.Millisecond * 10) + store.Must(ss.Team().RemoveMember(teamId, u1.Id)) store.Must(ss.Team().RemoveMember(teamId, u2.Id)) - store.Must(ss.User().UpdateUpdateAt(u1.Id)) - store.Must(ss.User().UpdateUpdateAt(u2.Id)) + u1.UpdateAt = store.Must(ss.User().UpdateUpdateAt(u1.Id)).(int64) + u2.UpdateAt = store.Must(ss.User().UpdateUpdateAt(u2.Id)).(int64) - if er3 := <-ss.User().GetEtagForProfilesNotInTeam(teamId); er3.Err != nil { - t.Fatal(er3.Err) - } else { - etag3 = er3.Data.(string) - t.Log(etag3) - if etag1 == etag3 || etag3 == etag2 { - t.Fatalf("etag should have changed") - } - } + t.Run("etag for profiles not in team 1 after second update", func(t *testing.T) { + result := <-ss.User().GetEtagForProfilesNotInTeam(teamId) + require.Nil(t, result.Err) + etag3 = result.Data.(string) + require.NotEqual(t, etag1, etag3, "etag should have changed") + require.NotEqual(t, etag2, etag3, "etag should have changed") + }) - if r3 := <-ss.User().GetProfilesNotInTeam(teamId, 0, 100000); r3.Err != nil { - t.Fatal(r3.Err) - } else { - users := r3.Data.([]*model.User) - found1, found2 := false, false - for _, u := range users { - if u.Id == u2.Id { - found2 = true - } - if u.Id == u1.Id { - found1 = true - } - } - - if !found1 || !found2 { - t.Fatal("missing user1 or user2") - } - } + t.Run("get not in team 1, offset 0, limit 100000 after second update", func(t *testing.T) { + result := <-ss.User().GetProfilesNotInTeam(teamId, 0, 100000) + require.Nil(t, result.Err) + assert.Equal(t, []*model.User{ + sanitized(u1), + sanitized(u2), + sanitized(u3), + }, result.Data.([]*model.User)) + }) + // Ensure update at timestamp changes time.Sleep(time.Millisecond * 10) - u3 := &model.User{} - u3.Email = MakeEmail() - store.Must(ss.User().Save(u3)) - defer func() { store.Must(ss.User().PermanentDelete(u3.Id)) }() - store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u3.Id}, -1)) - store.Must(ss.User().UpdateUpdateAt(u3.Id)) - if er4 := <-ss.User().GetEtagForProfilesNotInTeam(teamId); er4.Err != nil { - t.Fatal(er4.Err) - } else { - etag4 := er4.Data.(string) - t.Log(etag4) - if etag4 != etag3 { - t.Fatalf("etag should be the same") - } - } + u4 := &model.User{} + u4.Email = MakeEmail() + store.Must(ss.User().Save(u4)) + defer func() { store.Must(ss.User().PermanentDelete(u4.Id)) }() + store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: u4.Id}, -1)) + + t.Run("etag for profiles not in team 1 after addition to team", func(t *testing.T) { + result := <-ss.User().GetEtagForProfilesNotInTeam(teamId) + require.Nil(t, result.Err) + etag4 := result.Data.(string) + require.Equal(t, etag3, etag4, "etag should not have changed") + }) } func testUserStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) { @@ -2742,35 +3003,45 @@ func testUserStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) { } func testUserStoreGetAllAfter(t *testing.T, ss store.Store) { - u1 := model.User{ + u1 := store.Must(ss.User().Save(&model.User{ Email: MakeEmail(), Username: model.NewId(), Roles: "system_user system_admin system_post_all", - } - store.Must(ss.User().Save(&u1)) + })).(*model.User) defer func() { store.Must(ss.User().PermanentDelete(u1.Id)) }() - r1 := <-ss.User().GetAllAfter(10000, strings.Repeat("0", 26)) - require.Nil(t, r1.Err) + u2 := store.Must(ss.User().Save(&model.User{ + Email: MakeEmail(), + Username: "u2" + model.NewId(), + })).(*model.User) + defer func() { store.Must(ss.User().PermanentDelete(u2.Id)) }() - d1 := r1.Data.([]*model.User) - - found := false - for _, u := range d1 { - - if u.Id == u1.Id { - found = true - assert.Equal(t, u1.Id, u.Id) - assert.Equal(t, u1.Email, u.Email) - } + expected := []*model.User{u1, u2} + if strings.Compare(u2.Id, u1.Id) < 0 { + expected = []*model.User{u2, u1} } - assert.True(t, found) - r2 := <-ss.User().GetAllAfter(10000, u1.Id) - require.Nil(t, r2.Err) + t.Run("get after lowest possible id", func(t *testing.T) { + result := <-ss.User().GetAllAfter(10000, strings.Repeat("0", 26)) + require.Nil(t, result.Err) - d2 := r2.Data.([]*model.User) - for _, u := range d2 { - assert.NotEqual(t, u1.Id, u.Id) - } + actual := result.Data.([]*model.User) + assert.Equal(t, expected, actual) + }) + + t.Run("get after first user", func(t *testing.T) { + result := <-ss.User().GetAllAfter(10000, expected[0].Id) + require.Nil(t, result.Err) + + actual := result.Data.([]*model.User) + assert.Equal(t, []*model.User{expected[1]}, actual) + }) + + t.Run("get after second user", func(t *testing.T) { + result := <-ss.User().GetAllAfter(10000, expected[1].Id) + require.Nil(t, result.Err) + + actual := result.Data.([]*model.User) + assert.Equal(t, []*model.User{}, actual) + }) } diff --git a/vendor/github.com/Masterminds/squirrel/.gitignore b/vendor/github.com/Masterminds/squirrel/.gitignore new file mode 100644 index 0000000000..4a0699f0b7 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/.gitignore @@ -0,0 +1 @@ +squirrel.test \ No newline at end of file diff --git a/vendor/github.com/Masterminds/squirrel/.travis.yml b/vendor/github.com/Masterminds/squirrel/.travis.yml new file mode 100644 index 0000000000..06ee48eee0 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/.travis.yml @@ -0,0 +1,34 @@ +language: go + +go: + - 1.8.x + - 1.9.x + - 1.10.x + - 1.11.x + - tip + +services: + - mysql + - postgresql + +# Setting sudo access to false will let Travis CI use containers rather than +# VMs to run the tests. For more details see: +# - http://docs.travis-ci.com/user/workers/container-based-infrastructure/ +# - http://docs.travis-ci.com/user/workers/standard-infrastructure/ +sudo: false + +install: + - go get -t -tags integration + - go install github.com/mattn/go-sqlite3 # Precompile so test timing is accurate + +before_script: + - mysql -e 'CREATE DATABASE squirrel;' + - psql -c 'CREATE DATABASE squirrel;' -U postgres + +script: + - go test -tags integration -args -driver sqlite3 + - go test -tags integration -args -driver mysql -dataSource travis@/squirrel + - go test -tags integration -args -driver postgres -dataSource 'postgres://postgres@localhost/squirrel?sslmode=disable' + +notifications: + irc: "irc.freenode.net#masterminds" diff --git a/vendor/github.com/Masterminds/squirrel/LICENSE.txt b/vendor/github.com/Masterminds/squirrel/LICENSE.txt new file mode 100644 index 0000000000..74c20a2b97 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/LICENSE.txt @@ -0,0 +1,23 @@ +Squirrel +The Masterminds +Copyright (C) 2014-2015, Lann Martin +Copyright (C) 2015-2016, Google +Copyright (C) 2015, Matt Farina and Matt Butcher + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in +all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +THE SOFTWARE. diff --git a/vendor/github.com/Masterminds/squirrel/README.md b/vendor/github.com/Masterminds/squirrel/README.md new file mode 100644 index 0000000000..6ea9747c59 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/README.md @@ -0,0 +1,146 @@ +# Squirrel - fluent SQL generator for Go + +```go +import "gopkg.in/Masterminds/squirrel.v1" +``` +or if you prefer using `master` (which may be arbitrarily ahead of or behind `v1`): + +**NOTE:** as of Go 1.6, `go get` correctly clones the Github default branch (which is `v1` in this repo). +```go +import "github.com/Masterminds/squirrel" +``` + +[![GoDoc](https://godoc.org/github.com/Masterminds/squirrel?status.png)](https://godoc.org/github.com/Masterminds/squirrel) +[![Build Status](https://travis-ci.org/Masterminds/squirrel.svg?branch=v1)](https://travis-ci.org/Masterminds/squirrel) + +_**Note:** This project has moved from `github.com/lann/squirrel` to +`github.com/Masterminds/squirrel`. Lann remains the architect of the +project, but we're helping him curate. + +**Squirrel is not an ORM.** For an application of Squirrel, check out +[structable, a table-struct mapper](https://github.com/Masterminds/structable) + + +Squirrel helps you build SQL queries from composable parts: + +```go +import sq "github.com/Masterminds/squirrel" + +users := sq.Select("*").From("users").Join("emails USING (email_id)") + +active := users.Where(sq.Eq{"deleted_at": nil}) + +sql, args, err := active.ToSql() + +sql == "SELECT * FROM users JOIN emails USING (email_id) WHERE deleted_at IS NULL" +``` + +```go +sql, args, err := sq. + Insert("users").Columns("name", "age"). + Values("moe", 13).Values("larry", sq.Expr("? + 5", 12)). + ToSql() + +sql == "INSERT INTO users (name,age) VALUES (?,?),(?,? + 5)" +``` + +Squirrel can also execute queries directly: + +```go +stooges := users.Where(sq.Eq{"username": []string{"moe", "larry", "curly", "shemp"}}) +three_stooges := stooges.Limit(3) +rows, err := three_stooges.RunWith(db).Query() + +// Behaves like: +rows, err := db.Query("SELECT * FROM users WHERE username IN (?,?,?,?) LIMIT 3", + "moe", "larry", "curly", "shemp") +``` + +Squirrel makes conditional query building a breeze: + +```go +if len(q) > 0 { + users = users.Where("name LIKE ?", fmt.Sprint("%", q, "%")) +} +``` + +Squirrel wants to make your life easier: + +```go +// StmtCache caches Prepared Stmts for you +dbCache := sq.NewStmtCacher(db) + +// StatementBuilder keeps your syntax neat +mydb := sq.StatementBuilder.RunWith(dbCache) +select_users := mydb.Select("*").From("users") +``` + +Squirrel loves PostgreSQL: + +```go +psql := sq.StatementBuilder.PlaceholderFormat(sq.Dollar) + +// You use question marks for placeholders... +sql, _, _ := psql.Select("*").From("elephants").Where("name IN (?,?)", "Dumbo", "Verna").ToSql() + +/// ...squirrel replaces them using PlaceholderFormat. +sql == "SELECT * FROM elephants WHERE name IN ($1,$2)" + + +/// You can retrieve id ... +query := sq.Insert("nodes"). + Columns("uuid", "type", "data"). + Values(node.Uuid, node.Type, node.Data). + Suffix("RETURNING \"id\""). + RunWith(m.db). + PlaceholderFormat(sq.Dollar) + +query.QueryRow().Scan(&node.id) +``` + +You can escape question mask by inserting two question marks: + +```sql +SELECT * FROM nodes WHERE meta->'format' ??| array[?,?] +``` + +will generate with the Dollar Placeholder: + +```sql +SELECT * FROM nodes WHERE meta->'format' ?| array[$1,$2] +``` + +## FAQ + +* **How can I build an IN query on composite keys / tuples, e.g. `WHERE (col1, col2) IN ((1,2),(3,4))`? ([#104](https://github.com/Masterminds/squirrel/issues/104))** + + Squirrel does not explicitly support tuples, but you can get the same effect with e.g.: + + ```go + sq.Or{ + sq.Eq{"col1": 1, "col2": 2}, + sq.Eq{"col1": 3, "col2": 4}} + ``` + + ```sql + WHERE (col1 = 1 AND col2 = 2) OR (col1 = 3 AND col2 = 4) + ``` + + (which should produce the same query plan as the tuple version) + +* **Why doesn't `Eq{"mynumber": []uint8{1,2,3}}` turn into an `IN` query? ([#114](https://github.com/Masterminds/squirrel/issues/114))** + + Values of type `[]byte` are handled specially by `database/sql`. In Go, [`byte` is just an alias of `uint8`](https://golang.org/pkg/builtin/#byte), so there is no way to distinguish `[]uint8` from `[]byte`. + +* **Some features are poorly documented!** + + This isn't a frequent complaints section! + +* **Some features are poorly documented?** + + Yes. The tests should be considered a part of the documentation; take a look at those for ideas on how to express more complex queries. + +## License + +Squirrel is released under the +[MIT License](http://www.opensource.org/licenses/MIT). diff --git a/vendor/github.com/Masterminds/squirrel/case.go b/vendor/github.com/Masterminds/squirrel/case.go new file mode 100644 index 0000000000..2eb69dd5c9 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/case.go @@ -0,0 +1,118 @@ +package squirrel + +import ( + "bytes" + "errors" + + "github.com/lann/builder" +) + +func init() { + builder.Register(CaseBuilder{}, caseData{}) +} + +// sqlizerBuffer is a helper that allows to write many Sqlizers one by one +// without constant checks for errors that may come from Sqlizer +type sqlizerBuffer struct { + bytes.Buffer + args []interface{} + err error +} + +// WriteSql converts Sqlizer to SQL strings and writes it to buffer +func (b *sqlizerBuffer) WriteSql(item Sqlizer) { + if b.err != nil { + return + } + + var str string + var args []interface{} + str, args, b.err = item.ToSql() + + if b.err != nil { + return + } + + b.WriteString(str) + b.WriteByte(' ') + b.args = append(b.args, args...) +} + +func (b *sqlizerBuffer) ToSql() (string, []interface{}, error) { + return b.String(), b.args, b.err +} + +// whenPart is a helper structure to describe SQLs "WHEN ... THEN ..." expression +type whenPart struct { + when Sqlizer + then Sqlizer +} + +func newWhenPart(when interface{}, then interface{}) whenPart { + return whenPart{newPart(when), newPart(then)} +} + +// caseData holds all the data required to build a CASE SQL construct +type caseData struct { + What Sqlizer + WhenParts []whenPart + Else Sqlizer +} + +// ToSql implements Sqlizer +func (d *caseData) ToSql() (sqlStr string, args []interface{}, err error) { + if len(d.WhenParts) == 0 { + err = errors.New("case expression must contain at lease one WHEN clause") + + return + } + + sql := sqlizerBuffer{} + + sql.WriteString("CASE ") + if d.What != nil { + sql.WriteSql(d.What) + } + + for _, p := range d.WhenParts { + sql.WriteString("WHEN ") + sql.WriteSql(p.when) + sql.WriteString("THEN ") + sql.WriteSql(p.then) + } + + if d.Else != nil { + sql.WriteString("ELSE ") + sql.WriteSql(d.Else) + } + + sql.WriteString("END") + + return sql.ToSql() +} + +// CaseBuilder builds SQL CASE construct which could be used as parts of queries. +type CaseBuilder builder.Builder + +// ToSql builds the query into a SQL string and bound args. +func (b CaseBuilder) ToSql() (string, []interface{}, error) { + data := builder.GetStruct(b).(caseData) + return data.ToSql() +} + +// what sets optional value for CASE construct "CASE [value] ..." +func (b CaseBuilder) what(expr interface{}) CaseBuilder { + return builder.Set(b, "What", newPart(expr)).(CaseBuilder) +} + +// When adds "WHEN ... THEN ..." part to CASE construct +func (b CaseBuilder) When(when interface{}, then interface{}) CaseBuilder { + // TODO: performance hint: replace slice of WhenPart with just slice of parts + // where even indices of the slice belong to "when"s and odd indices belong to "then"s + return builder.Append(b, "WhenParts", newWhenPart(when, then)).(CaseBuilder) +} + +// What sets optional "ELSE ..." part for CASE construct +func (b CaseBuilder) Else(expr interface{}) CaseBuilder { + return builder.Set(b, "Else", newPart(expr)).(CaseBuilder) +} diff --git a/vendor/github.com/Masterminds/squirrel/delete.go b/vendor/github.com/Masterminds/squirrel/delete.go new file mode 100644 index 0000000000..41aebbbe5d --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/delete.go @@ -0,0 +1,164 @@ +package squirrel + +import ( + "bytes" + "database/sql" + "fmt" + "strings" + + "github.com/lann/builder" +) + +type deleteData struct { + PlaceholderFormat PlaceholderFormat + RunWith BaseRunner + Prefixes exprs + From string + WhereParts []Sqlizer + OrderBys []string + Limit string + Offset string + Suffixes exprs +} + +func (d *deleteData) Exec() (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return ExecWith(d.RunWith, d) +} + +func (d *deleteData) ToSql() (sqlStr string, args []interface{}, err error) { + if len(d.From) == 0 { + err = fmt.Errorf("delete statements must specify a From table") + return + } + + sql := &bytes.Buffer{} + + if len(d.Prefixes) > 0 { + args, _ = d.Prefixes.AppendToSql(sql, " ", args) + sql.WriteString(" ") + } + + sql.WriteString("DELETE FROM ") + sql.WriteString(d.From) + + if len(d.WhereParts) > 0 { + sql.WriteString(" WHERE ") + args, err = appendToSql(d.WhereParts, sql, " AND ", args) + if err != nil { + return + } + } + + if len(d.OrderBys) > 0 { + sql.WriteString(" ORDER BY ") + sql.WriteString(strings.Join(d.OrderBys, ", ")) + } + + if len(d.Limit) > 0 { + sql.WriteString(" LIMIT ") + sql.WriteString(d.Limit) + } + + if len(d.Offset) > 0 { + sql.WriteString(" OFFSET ") + sql.WriteString(d.Offset) + } + + if len(d.Suffixes) > 0 { + sql.WriteString(" ") + args, _ = d.Suffixes.AppendToSql(sql, " ", args) + } + + sqlStr, err = d.PlaceholderFormat.ReplacePlaceholders(sql.String()) + return +} + +// Builder + +// DeleteBuilder builds SQL DELETE statements. +type DeleteBuilder builder.Builder + +func init() { + builder.Register(DeleteBuilder{}, deleteData{}) +} + +// Format methods + +// PlaceholderFormat sets PlaceholderFormat (e.g. Question or Dollar) for the +// query. +func (b DeleteBuilder) PlaceholderFormat(f PlaceholderFormat) DeleteBuilder { + return builder.Set(b, "PlaceholderFormat", f).(DeleteBuilder) +} + +// Runner methods + +// RunWith sets a Runner (like database/sql.DB) to be used with e.g. Exec. +func (b DeleteBuilder) RunWith(runner BaseRunner) DeleteBuilder { + return setRunWith(b, runner).(DeleteBuilder) +} + +// Exec builds and Execs the query with the Runner set by RunWith. +func (b DeleteBuilder) Exec() (sql.Result, error) { + data := builder.GetStruct(b).(deleteData) + return data.Exec() +} + +// SQL methods + +// ToSql builds the query into a SQL string and bound args. +func (b DeleteBuilder) ToSql() (string, []interface{}, error) { + data := builder.GetStruct(b).(deleteData) + return data.ToSql() +} + +// Prefix adds an expression to the beginning of the query +func (b DeleteBuilder) Prefix(sql string, args ...interface{}) DeleteBuilder { + return builder.Append(b, "Prefixes", Expr(sql, args...)).(DeleteBuilder) +} + +// From sets the table to be deleted from. +func (b DeleteBuilder) From(from string) DeleteBuilder { + return builder.Set(b, "From", from).(DeleteBuilder) +} + +// Where adds WHERE expressions to the query. +// +// See SelectBuilder.Where for more information. +func (b DeleteBuilder) Where(pred interface{}, args ...interface{}) DeleteBuilder { + return builder.Append(b, "WhereParts", newWherePart(pred, args...)).(DeleteBuilder) +} + +// OrderBy adds ORDER BY expressions to the query. +func (b DeleteBuilder) OrderBy(orderBys ...string) DeleteBuilder { + return builder.Extend(b, "OrderBys", orderBys).(DeleteBuilder) +} + +// Limit sets a LIMIT clause on the query. +func (b DeleteBuilder) Limit(limit uint64) DeleteBuilder { + return builder.Set(b, "Limit", fmt.Sprintf("%d", limit)).(DeleteBuilder) +} + +// Offset sets a OFFSET clause on the query. +func (b DeleteBuilder) Offset(offset uint64) DeleteBuilder { + return builder.Set(b, "Offset", fmt.Sprintf("%d", offset)).(DeleteBuilder) +} + +// Suffix adds an expression to the end of the query +func (b DeleteBuilder) Suffix(sql string, args ...interface{}) DeleteBuilder { + return builder.Append(b, "Suffixes", Expr(sql, args...)).(DeleteBuilder) +} + +func (b DeleteBuilder) Query() (*sql.Rows, error) { + data := builder.GetStruct(b).(deleteData) + return data.Query() +} + +func (d *deleteData) Query() (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return QueryWith(d.RunWith, d) +} diff --git a/vendor/github.com/Masterminds/squirrel/delete_ctx.go b/vendor/github.com/Masterminds/squirrel/delete_ctx.go new file mode 100644 index 0000000000..ecdf7ef03f --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/delete_ctx.go @@ -0,0 +1,27 @@ +// +build go1.8 + +package squirrel + +import ( + "context" + "database/sql" + + "github.com/lann/builder" +) + +func (d *deleteData) ExecContext(ctx context.Context) (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(ExecerContext) + if !ok { + return nil, NoContextSupport + } + return ExecContextWith(ctx, ctxRunner, d) +} + +// ExecContext builds and ExecContexts the query with the Runner set by RunWith. +func (b DeleteBuilder) ExecContext(ctx context.Context) (sql.Result, error) { + data := builder.GetStruct(b).(deleteData) + return data.ExecContext(ctx) +} diff --git a/vendor/github.com/Masterminds/squirrel/expr.go b/vendor/github.com/Masterminds/squirrel/expr.go new file mode 100644 index 0000000000..cfb7521646 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/expr.go @@ -0,0 +1,351 @@ +package squirrel + +import ( + "database/sql/driver" + "fmt" + "io" + "reflect" + "sort" + "strings" +) + +const ( + // Portable true/false literals. + sqlTrue = "(1=1)" + sqlFalse = "(1=0)" +) + +type expr struct { + sql string + args []interface{} +} + +// Expr builds value expressions for InsertBuilder and UpdateBuilder. +// +// Ex: +// .Values(Expr("FROM_UNIXTIME(?)", t)) +func Expr(sql string, args ...interface{}) expr { + return expr{sql: sql, args: args} +} + +func (e expr) ToSql() (sql string, args []interface{}, err error) { + return e.sql, e.args, nil +} + +type exprs []expr + +func (es exprs) AppendToSql(w io.Writer, sep string, args []interface{}) ([]interface{}, error) { + for i, e := range es { + if i > 0 { + _, err := io.WriteString(w, sep) + if err != nil { + return nil, err + } + } + _, err := io.WriteString(w, e.sql) + if err != nil { + return nil, err + } + args = append(args, e.args...) + } + return args, nil +} + +// aliasExpr helps to alias part of SQL query generated with underlying "expr" +type aliasExpr struct { + expr Sqlizer + alias string +} + +// Alias allows to define alias for column in SelectBuilder. Useful when column is +// defined as complex expression like IF or CASE +// Ex: +// .Column(Alias(caseStmt, "case_column")) +func Alias(expr Sqlizer, alias string) aliasExpr { + return aliasExpr{expr, alias} +} + +func (e aliasExpr) ToSql() (sql string, args []interface{}, err error) { + sql, args, err = e.expr.ToSql() + if err == nil { + sql = fmt.Sprintf("(%s) AS %s", sql, e.alias) + } + return +} + +// Eq is syntactic sugar for use with Where/Having/Set methods. +// Ex: +// .Where(Eq{"id": 1}) +type Eq map[string]interface{} + +func (eq Eq) toSQL(useNotOpr bool) (sql string, args []interface{}, err error) { + if len(eq) == 0 { + // Empty Sql{} evaluates to true. + sql = sqlTrue + return + } + + var ( + exprs []string + equalOpr = "=" + inOpr = "IN" + nullOpr = "IS" + inEmptyExpr = sqlFalse + ) + + if useNotOpr { + equalOpr = "<>" + inOpr = "NOT IN" + nullOpr = "IS NOT" + inEmptyExpr = sqlTrue + } + + sortedKeys := getSortedKeys(eq) + for _, key := range sortedKeys { + var expr string + val := eq[key] + + switch v := val.(type) { + case driver.Valuer: + if val, err = v.Value(); err != nil { + return + } + } + + r := reflect.ValueOf(val) + if r.Kind() == reflect.Ptr { + if r.IsNil() { + val = nil + } else { + val = r.Elem().Interface() + } + } + + if val == nil { + expr = fmt.Sprintf("%s %s NULL", key, nullOpr) + } else { + if isListType(val) { + valVal := reflect.ValueOf(val) + if valVal.Len() == 0 { + expr = inEmptyExpr + if args == nil { + args = []interface{}{} + } + } else { + for i := 0; i < valVal.Len(); i++ { + args = append(args, valVal.Index(i).Interface()) + } + expr = fmt.Sprintf("%s %s (%s)", key, inOpr, Placeholders(valVal.Len())) + } + } else { + expr = fmt.Sprintf("%s %s ?", key, equalOpr) + args = append(args, val) + } + } + exprs = append(exprs, expr) + } + sql = strings.Join(exprs, " AND ") + return +} + +func (eq Eq) ToSql() (sql string, args []interface{}, err error) { + return eq.toSQL(false) +} + +// NotEq is syntactic sugar for use with Where/Having/Set methods. +// Ex: +// .Where(NotEq{"id": 1}) == "id <> 1" +type NotEq Eq + +func (neq NotEq) ToSql() (sql string, args []interface{}, err error) { + return Eq(neq).toSQL(true) +} + +// Like is syntactic sugar for use with LIKE conditions. +// Ex: +// .Where(Like{"name": "%irrel"}) +type Like map[string]interface{} + +func (lk Like) toSql(opposite bool) (sql string, args []interface{}, err error) { + var ( + exprs []string + opr = "LIKE" + ) + + if opposite { + opr = "NOT LIKE" + } + + for key, val := range lk { + expr := "" + + switch v := val.(type) { + case driver.Valuer: + if val, err = v.Value(); err != nil { + return + } + } + + if val == nil { + err = fmt.Errorf("cannot use null with like operators") + return + } else { + if isListType(val) { + err = fmt.Errorf("cannot use array or slice with like operators") + return + } else { + expr = fmt.Sprintf("%s %s ?", key, opr) + args = append(args, val) + } + } + exprs = append(exprs, expr) + } + sql = strings.Join(exprs, " AND ") + return +} + +func (lk Like) ToSql() (sql string, args []interface{}, err error) { + return lk.toSql(false) +} + +// NotLike is syntactic sugar for use with LIKE conditions. +// Ex: +// .Where(NotLike{"name": "%irrel"}) +type NotLike Like + +func (nlk NotLike) ToSql() (sql string, args []interface{}, err error) { + return Like(nlk).toSql(true) +} + +// Lt is syntactic sugar for use with Where/Having/Set methods. +// Ex: +// .Where(Lt{"id": 1}) +type Lt map[string]interface{} + +func (lt Lt) toSql(opposite, orEq bool) (sql string, args []interface{}, err error) { + var ( + exprs []string + opr = "<" + ) + + if opposite { + opr = ">" + } + + if orEq { + opr = fmt.Sprintf("%s%s", opr, "=") + } + + sortedKeys := getSortedKeys(lt) + for _, key := range sortedKeys { + var expr string + val := lt[key] + + switch v := val.(type) { + case driver.Valuer: + if val, err = v.Value(); err != nil { + return + } + } + + if val == nil { + err = fmt.Errorf("cannot use null with less than or greater than operators") + return + } + if isListType(val) { + err = fmt.Errorf("cannot use array or slice with less than or greater than operators") + return + } + expr = fmt.Sprintf("%s %s ?", key, opr) + args = append(args, val) + + exprs = append(exprs, expr) + } + sql = strings.Join(exprs, " AND ") + return +} + +func (lt Lt) ToSql() (sql string, args []interface{}, err error) { + return lt.toSql(false, false) +} + +// LtOrEq is syntactic sugar for use with Where/Having/Set methods. +// Ex: +// .Where(LtOrEq{"id": 1}) == "id <= 1" +type LtOrEq Lt + +func (ltOrEq LtOrEq) ToSql() (sql string, args []interface{}, err error) { + return Lt(ltOrEq).toSql(false, true) +} + +// Gt is syntactic sugar for use with Where/Having/Set methods. +// Ex: +// .Where(Gt{"id": 1}) == "id > 1" +type Gt Lt + +func (gt Gt) ToSql() (sql string, args []interface{}, err error) { + return Lt(gt).toSql(true, false) +} + +// GtOrEq is syntactic sugar for use with Where/Having/Set methods. +// Ex: +// .Where(GtOrEq{"id": 1}) == "id >= 1" +type GtOrEq Lt + +func (gtOrEq GtOrEq) ToSql() (sql string, args []interface{}, err error) { + return Lt(gtOrEq).toSql(true, true) +} + +type conj []Sqlizer + +func (c conj) join(sep, defaultExpr string) (sql string, args []interface{}, err error) { + if len(c) == 0 { + return defaultExpr, []interface{}{}, nil + } + var sqlParts []string + for _, sqlizer := range c { + partSQL, partArgs, err := sqlizer.ToSql() + if err != nil { + return "", nil, err + } + if partSQL != "" { + sqlParts = append(sqlParts, partSQL) + args = append(args, partArgs...) + } + } + if len(sqlParts) > 0 { + sql = fmt.Sprintf("(%s)", strings.Join(sqlParts, sep)) + } + return +} + +// And conjunction Sqlizers +type And conj + +func (a And) ToSql() (string, []interface{}, error) { + return conj(a).join(" AND ", sqlTrue) +} + +// Or conjunction Sqlizers +type Or conj + +func (o Or) ToSql() (string, []interface{}, error) { + return conj(o).join(" OR ", sqlFalse) +} + +func getSortedKeys(exp map[string]interface{}) []string { + sortedKeys := make([]string, 0, len(exp)) + for k := range exp { + sortedKeys = append(sortedKeys, k) + } + sort.Strings(sortedKeys) + return sortedKeys +} + +func isListType(val interface{}) bool { + if driver.IsValue(val) { + return false + } + valVal := reflect.ValueOf(val) + return valVal.Kind() == reflect.Array || valVal.Kind() == reflect.Slice +} diff --git a/vendor/github.com/Masterminds/squirrel/go.mod b/vendor/github.com/Masterminds/squirrel/go.mod new file mode 100644 index 0000000000..ba46589a6f --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/go.mod @@ -0,0 +1,9 @@ +module github.com/Masterminds/squirrel + +require ( + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/lann/builder v0.0.0-20180802200727-47ae307949d0 + github.com/lann/ps v0.0.0-20150810152359-62de8c46ede0 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/stretchr/testify v1.2.2 +) diff --git a/vendor/github.com/Masterminds/squirrel/go.sum b/vendor/github.com/Masterminds/squirrel/go.sum new file mode 100644 index 0000000000..a92f768e60 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/go.sum @@ -0,0 +1,10 @@ +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/lann/builder v0.0.0-20180802200727-47ae307949d0 h1:SOEGU9fKiNWd/HOJuq6+3iTQz8KNCLtVX6idSoTLdUw= +github.com/lann/builder v0.0.0-20180802200727-47ae307949d0/go.mod h1:dXGbAdH5GtBTC4WfIxhKZfyBF/HBFgRZSWwZ9g/He9o= +github.com/lann/ps v0.0.0-20150810152359-62de8c46ede0 h1:P6pPBnrTSX3DEVR4fDembhRWSsG5rVo6hYhAB/ADZrk= +github.com/lann/ps v0.0.0-20150810152359-62de8c46ede0/go.mod h1:vmVJ0l/dxyfGW6FmdpVm2joNMFikkuWg0EoCKLGUMNw= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/testify v1.2.2 h1:bSDNvY7ZPG5RlJ8otE/7V6gMiyenm9RtJ7IUVIAoJ1w= +github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= diff --git a/vendor/github.com/Masterminds/squirrel/insert.go b/vendor/github.com/Masterminds/squirrel/insert.go new file mode 100644 index 0000000000..becf81e40d --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/insert.go @@ -0,0 +1,258 @@ +package squirrel + +import ( + "bytes" + "database/sql" + "errors" + "fmt" + "io" + "sort" + "strings" + + "github.com/lann/builder" +) + +type insertData struct { + PlaceholderFormat PlaceholderFormat + RunWith BaseRunner + Prefixes exprs + Options []string + Into string + Columns []string + Values [][]interface{} + Suffixes exprs + Select *SelectBuilder +} + +func (d *insertData) Exec() (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return ExecWith(d.RunWith, d) +} + +func (d *insertData) Query() (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return QueryWith(d.RunWith, d) +} + +func (d *insertData) QueryRow() RowScanner { + if d.RunWith == nil { + return &Row{err: RunnerNotSet} + } + queryRower, ok := d.RunWith.(QueryRower) + if !ok { + return &Row{err: RunnerNotQueryRunner} + } + return QueryRowWith(queryRower, d) +} + +func (d *insertData) ToSql() (sqlStr string, args []interface{}, err error) { + if len(d.Into) == 0 { + err = errors.New("insert statements must specify a table") + return + } + if len(d.Values) == 0 && d.Select == nil { + err = errors.New("insert statements must have at least one set of values or select clause") + return + } + + sql := &bytes.Buffer{} + + if len(d.Prefixes) > 0 { + args, _ = d.Prefixes.AppendToSql(sql, " ", args) + sql.WriteString(" ") + } + + sql.WriteString("INSERT ") + + if len(d.Options) > 0 { + sql.WriteString(strings.Join(d.Options, " ")) + sql.WriteString(" ") + } + + sql.WriteString("INTO ") + sql.WriteString(d.Into) + sql.WriteString(" ") + + if len(d.Columns) > 0 { + sql.WriteString("(") + sql.WriteString(strings.Join(d.Columns, ",")) + sql.WriteString(") ") + } + + if d.Select != nil { + args, err = d.appendSelectToSQL(sql, args) + } else { + args, err = d.appendValuesToSQL(sql, args) + } + if err != nil { + return + } + + if len(d.Suffixes) > 0 { + sql.WriteString(" ") + args, _ = d.Suffixes.AppendToSql(sql, " ", args) + } + + sqlStr, err = d.PlaceholderFormat.ReplacePlaceholders(sql.String()) + return +} + +func (d *insertData) appendValuesToSQL(w io.Writer, args []interface{}) ([]interface{}, error) { + if len(d.Values) == 0 { + return args, errors.New("values for insert statements are not set") + } + + io.WriteString(w, "VALUES ") + + valuesStrings := make([]string, len(d.Values)) + for r, row := range d.Values { + valueStrings := make([]string, len(row)) + for v, val := range row { + e, isExpr := val.(expr) + if isExpr { + valueStrings[v] = e.sql + args = append(args, e.args...) + } else { + valueStrings[v] = "?" + args = append(args, val) + } + } + valuesStrings[r] = fmt.Sprintf("(%s)", strings.Join(valueStrings, ",")) + } + + io.WriteString(w, strings.Join(valuesStrings, ",")) + + return args, nil +} + +func (d *insertData) appendSelectToSQL(w io.Writer, args []interface{}) ([]interface{}, error) { + if d.Select == nil { + return args, errors.New("select clause for insert statements are not set") + } + + selectClause, sArgs, err := d.Select.ToSql() + if err != nil { + return args, err + } + + io.WriteString(w, selectClause) + args = append(args, sArgs...) + + return args, nil +} + +// Builder + +// InsertBuilder builds SQL INSERT statements. +type InsertBuilder builder.Builder + +func init() { + builder.Register(InsertBuilder{}, insertData{}) +} + +// Format methods + +// PlaceholderFormat sets PlaceholderFormat (e.g. Question or Dollar) for the +// query. +func (b InsertBuilder) PlaceholderFormat(f PlaceholderFormat) InsertBuilder { + return builder.Set(b, "PlaceholderFormat", f).(InsertBuilder) +} + +// Runner methods + +// RunWith sets a Runner (like database/sql.DB) to be used with e.g. Exec. +func (b InsertBuilder) RunWith(runner BaseRunner) InsertBuilder { + return setRunWith(b, runner).(InsertBuilder) +} + +// Exec builds and Execs the query with the Runner set by RunWith. +func (b InsertBuilder) Exec() (sql.Result, error) { + data := builder.GetStruct(b).(insertData) + return data.Exec() +} + +// Query builds and Querys the query with the Runner set by RunWith. +func (b InsertBuilder) Query() (*sql.Rows, error) { + data := builder.GetStruct(b).(insertData) + return data.Query() +} + +// QueryRow builds and QueryRows the query with the Runner set by RunWith. +func (b InsertBuilder) QueryRow() RowScanner { + data := builder.GetStruct(b).(insertData) + return data.QueryRow() +} + +// Scan is a shortcut for QueryRow().Scan. +func (b InsertBuilder) Scan(dest ...interface{}) error { + return b.QueryRow().Scan(dest...) +} + +// SQL methods + +// ToSql builds the query into a SQL string and bound args. +func (b InsertBuilder) ToSql() (string, []interface{}, error) { + data := builder.GetStruct(b).(insertData) + return data.ToSql() +} + +// Prefix adds an expression to the beginning of the query +func (b InsertBuilder) Prefix(sql string, args ...interface{}) InsertBuilder { + return builder.Append(b, "Prefixes", Expr(sql, args...)).(InsertBuilder) +} + +// Options adds keyword options before the INTO clause of the query. +func (b InsertBuilder) Options(options ...string) InsertBuilder { + return builder.Extend(b, "Options", options).(InsertBuilder) +} + +// Into sets the INTO clause of the query. +func (b InsertBuilder) Into(from string) InsertBuilder { + return builder.Set(b, "Into", from).(InsertBuilder) +} + +// Columns adds insert columns to the query. +func (b InsertBuilder) Columns(columns ...string) InsertBuilder { + return builder.Extend(b, "Columns", columns).(InsertBuilder) +} + +// Values adds a single row's values to the query. +func (b InsertBuilder) Values(values ...interface{}) InsertBuilder { + return builder.Append(b, "Values", values).(InsertBuilder) +} + +// Suffix adds an expression to the end of the query +func (b InsertBuilder) Suffix(sql string, args ...interface{}) InsertBuilder { + return builder.Append(b, "Suffixes", Expr(sql, args...)).(InsertBuilder) +} + +// SetMap set columns and values for insert builder from a map of column name and value +// note that it will reset all previous columns and values was set if any +func (b InsertBuilder) SetMap(clauses map[string]interface{}) InsertBuilder { + // Keep the columns in a consistent order by sorting the column key string. + cols := make([]string, 0, len(clauses)) + for col := range clauses { + cols = append(cols, col) + } + sort.Strings(cols) + + vals := make([]interface{}, 0, len(clauses)) + for _, col := range cols { + vals = append(vals, clauses[col]) + } + + b = builder.Set(b, "Columns", cols).(InsertBuilder) + b = builder.Set(b, "Values", [][]interface{}{vals}).(InsertBuilder) + + return b +} + +// Select set Select clause for insert query +// If Values and Select are used, then Select has higher priority +func (b InsertBuilder) Select(sb SelectBuilder) InsertBuilder { + return builder.Set(b, "Select", &sb).(InsertBuilder) +} diff --git a/vendor/github.com/Masterminds/squirrel/insert_ctx.go b/vendor/github.com/Masterminds/squirrel/insert_ctx.go new file mode 100644 index 0000000000..4541c2fed3 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/insert_ctx.go @@ -0,0 +1,69 @@ +// +build go1.8 + +package squirrel + +import ( + "context" + "database/sql" + + "github.com/lann/builder" +) + +func (d *insertData) ExecContext(ctx context.Context) (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(ExecerContext) + if !ok { + return nil, NoContextSupport + } + return ExecContextWith(ctx, ctxRunner, d) +} + +func (d *insertData) QueryContext(ctx context.Context) (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(QueryerContext) + if !ok { + return nil, NoContextSupport + } + return QueryContextWith(ctx, ctxRunner, d) +} + +func (d *insertData) QueryRowContext(ctx context.Context) RowScanner { + if d.RunWith == nil { + return &Row{err: RunnerNotSet} + } + queryRower, ok := d.RunWith.(QueryRowerContext) + if !ok { + if _, ok := d.RunWith.(QueryerContext); !ok { + return &Row{err: RunnerNotQueryRunner} + } + return &Row{err: NoContextSupport} + } + return QueryRowContextWith(ctx, queryRower, d) +} + +// ExecContext builds and ExecContexts the query with the Runner set by RunWith. +func (b InsertBuilder) ExecContext(ctx context.Context) (sql.Result, error) { + data := builder.GetStruct(b).(insertData) + return data.ExecContext(ctx) +} + +// QueryContext builds and QueryContexts the query with the Runner set by RunWith. +func (b InsertBuilder) QueryContext(ctx context.Context) (*sql.Rows, error) { + data := builder.GetStruct(b).(insertData) + return data.QueryContext(ctx) +} + +// QueryRowContext builds and QueryRowContexts the query with the Runner set by RunWith. +func (b InsertBuilder) QueryRowContext(ctx context.Context) RowScanner { + data := builder.GetStruct(b).(insertData) + return data.QueryRowContext(ctx) +} + +// ScanContext is a shortcut for QueryRowContext().Scan. +func (b InsertBuilder) ScanContext(ctx context.Context, dest ...interface{}) error { + return b.QueryRowContext(ctx).Scan(dest...) +} diff --git a/vendor/github.com/Masterminds/squirrel/part.go b/vendor/github.com/Masterminds/squirrel/part.go new file mode 100644 index 0000000000..2926d03151 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/part.go @@ -0,0 +1,55 @@ +package squirrel + +import ( + "fmt" + "io" +) + +type part struct { + pred interface{} + args []interface{} +} + +func newPart(pred interface{}, args ...interface{}) Sqlizer { + return &part{pred, args} +} + +func (p part) ToSql() (sql string, args []interface{}, err error) { + switch pred := p.pred.(type) { + case nil: + // no-op + case Sqlizer: + sql, args, err = pred.ToSql() + case string: + sql = pred + args = p.args + default: + err = fmt.Errorf("expected string or Sqlizer, not %T", pred) + } + return +} + +func appendToSql(parts []Sqlizer, w io.Writer, sep string, args []interface{}) ([]interface{}, error) { + for i, p := range parts { + partSql, partArgs, err := p.ToSql() + if err != nil { + return nil, err + } else if len(partSql) == 0 { + continue + } + + if i > 0 { + _, err := io.WriteString(w, sep) + if err != nil { + return nil, err + } + } + + _, err = io.WriteString(w, partSql) + if err != nil { + return nil, err + } + args = append(args, partArgs...) + } + return args, nil +} diff --git a/vendor/github.com/Masterminds/squirrel/placeholder.go b/vendor/github.com/Masterminds/squirrel/placeholder.go new file mode 100644 index 0000000000..a016182394 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/placeholder.go @@ -0,0 +1,84 @@ +package squirrel + +import ( + "bytes" + "fmt" + "strings" +) + +// PlaceholderFormat is the interface that wraps the ReplacePlaceholders method. +// +// ReplacePlaceholders takes a SQL statement and replaces each question mark +// placeholder with a (possibly different) SQL placeholder. +type PlaceholderFormat interface { + ReplacePlaceholders(sql string) (string, error) +} + +var ( + // Question is a PlaceholderFormat instance that leaves placeholders as + // question marks. + Question = questionFormat{} + + // Dollar is a PlaceholderFormat instance that replaces placeholders with + // dollar-prefixed positional placeholders (e.g. $1, $2, $3). + Dollar = dollarFormat{} + + // Colon is a PlaceholderFormat instance that replaces placeholders with + // colon-prefixed positional placeholders (e.g. :1, :2, :3). + Colon = colonFormat{} +) + +type questionFormat struct{} + +func (questionFormat) ReplacePlaceholders(sql string) (string, error) { + return sql, nil +} + +type dollarFormat struct{} + +func (dollarFormat) ReplacePlaceholders(sql string) (string, error) { + return replacePositionalPlaceholders(sql, "$") +} + +type colonFormat struct{} + +func (colonFormat) ReplacePlaceholders(sql string) (string, error) { + return replacePositionalPlaceholders(sql, ":") +} + +// Placeholders returns a string with count ? placeholders joined with commas. +func Placeholders(count int) string { + if count < 1 { + return "" + } + + return strings.Repeat(",?", count)[1:] +} + +func replacePositionalPlaceholders(sql, prefix string) (string, error) { + buf := &bytes.Buffer{} + i := 0 + for { + p := strings.Index(sql, "?") + if p == -1 { + break + } + + if len(sql[p:]) > 1 && sql[p:p+2] == "??" { // escape ?? => ? + buf.WriteString(sql[:p]) + buf.WriteString("?") + if len(sql[p:]) == 1 { + break + } + sql = sql[p+2:] + } else { + i++ + buf.WriteString(sql[:p]) + fmt.Fprintf(buf, "%s%d", prefix, i) + sql = sql[p+1:] + } + } + + buf.WriteString(sql) + return buf.String(), nil +} diff --git a/vendor/github.com/Masterminds/squirrel/row.go b/vendor/github.com/Masterminds/squirrel/row.go new file mode 100644 index 0000000000..74ffda92bd --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/row.go @@ -0,0 +1,22 @@ +package squirrel + +// RowScanner is the interface that wraps the Scan method. +// +// Scan behaves like database/sql.Row.Scan. +type RowScanner interface { + Scan(...interface{}) error +} + +// Row wraps database/sql.Row to let squirrel return new errors on Scan. +type Row struct { + RowScanner + err error +} + +// Scan returns Row.err or calls RowScanner.Scan. +func (r *Row) Scan(dest ...interface{}) error { + if r.err != nil { + return r.err + } + return r.RowScanner.Scan(dest...) +} diff --git a/vendor/github.com/Masterminds/squirrel/select.go b/vendor/github.com/Masterminds/squirrel/select.go new file mode 100644 index 0000000000..0ec4dc6ba3 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/select.go @@ -0,0 +1,348 @@ +package squirrel + +import ( + "bytes" + "database/sql" + "fmt" + "strings" + + "github.com/lann/builder" +) + +type selectData struct { + PlaceholderFormat PlaceholderFormat + RunWith BaseRunner + Prefixes exprs + Options []string + Columns []Sqlizer + From Sqlizer + Joins []Sqlizer + WhereParts []Sqlizer + GroupBys []string + HavingParts []Sqlizer + OrderBys []string + Limit string + Offset string + Suffixes exprs +} + +func (d *selectData) Exec() (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return ExecWith(d.RunWith, d) +} + +func (d *selectData) Query() (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return QueryWith(d.RunWith, d) +} + +func (d *selectData) QueryRow() RowScanner { + if d.RunWith == nil { + return &Row{err: RunnerNotSet} + } + queryRower, ok := d.RunWith.(QueryRower) + if !ok { + return &Row{err: RunnerNotQueryRunner} + } + return QueryRowWith(queryRower, d) +} + +func (d *selectData) ToSql() (sqlStr string, args []interface{}, err error) { + sqlStr, args, err = d.toSql() + if err != nil { + return + } + + sqlStr, err = d.PlaceholderFormat.ReplacePlaceholders(sqlStr) + return +} + +func (d *selectData) toSqlRaw() (sqlStr string, args []interface{}, err error) { + return d.toSql() +} + +func (d *selectData) toSql() (sqlStr string, args []interface{}, err error) { + if len(d.Columns) == 0 { + err = fmt.Errorf("select statements must have at least one result column") + return + } + + sql := &bytes.Buffer{} + + if len(d.Prefixes) > 0 { + args, _ = d.Prefixes.AppendToSql(sql, " ", args) + sql.WriteString(" ") + } + + sql.WriteString("SELECT ") + + if len(d.Options) > 0 { + sql.WriteString(strings.Join(d.Options, " ")) + sql.WriteString(" ") + } + + if len(d.Columns) > 0 { + args, err = appendToSql(d.Columns, sql, ", ", args) + if err != nil { + return + } + } + + if d.From != nil { + sql.WriteString(" FROM ") + args, err = appendToSql([]Sqlizer{d.From}, sql, "", args) + if err != nil { + return + } + } + + if len(d.Joins) > 0 { + sql.WriteString(" ") + args, err = appendToSql(d.Joins, sql, " ", args) + if err != nil { + return + } + } + + if len(d.WhereParts) > 0 { + sql.WriteString(" WHERE ") + args, err = appendToSql(d.WhereParts, sql, " AND ", args) + if err != nil { + return + } + } + + if len(d.GroupBys) > 0 { + sql.WriteString(" GROUP BY ") + sql.WriteString(strings.Join(d.GroupBys, ", ")) + } + + if len(d.HavingParts) > 0 { + sql.WriteString(" HAVING ") + args, err = appendToSql(d.HavingParts, sql, " AND ", args) + if err != nil { + return + } + } + + if len(d.OrderBys) > 0 { + sql.WriteString(" ORDER BY ") + sql.WriteString(strings.Join(d.OrderBys, ", ")) + } + + if len(d.Limit) > 0 { + sql.WriteString(" LIMIT ") + sql.WriteString(d.Limit) + } + + if len(d.Offset) > 0 { + sql.WriteString(" OFFSET ") + sql.WriteString(d.Offset) + } + + if len(d.Suffixes) > 0 { + sql.WriteString(" ") + args, _ = d.Suffixes.AppendToSql(sql, " ", args) + } + + sqlStr = sql.String() + return +} + +// Builder + +// SelectBuilder builds SQL SELECT statements. +type SelectBuilder builder.Builder + +func init() { + builder.Register(SelectBuilder{}, selectData{}) +} + +// Format methods + +// PlaceholderFormat sets PlaceholderFormat (e.g. Question or Dollar) for the +// query. +func (b SelectBuilder) PlaceholderFormat(f PlaceholderFormat) SelectBuilder { + return builder.Set(b, "PlaceholderFormat", f).(SelectBuilder) +} + +// Runner methods + +// RunWith sets a Runner (like database/sql.DB) to be used with e.g. Exec. +func (b SelectBuilder) RunWith(runner BaseRunner) SelectBuilder { + return setRunWith(b, runner).(SelectBuilder) +} + +// Exec builds and Execs the query with the Runner set by RunWith. +func (b SelectBuilder) Exec() (sql.Result, error) { + data := builder.GetStruct(b).(selectData) + return data.Exec() +} + +// Query builds and Querys the query with the Runner set by RunWith. +func (b SelectBuilder) Query() (*sql.Rows, error) { + data := builder.GetStruct(b).(selectData) + return data.Query() +} + +// QueryRow builds and QueryRows the query with the Runner set by RunWith. +func (b SelectBuilder) QueryRow() RowScanner { + data := builder.GetStruct(b).(selectData) + return data.QueryRow() +} + +// Scan is a shortcut for QueryRow().Scan. +func (b SelectBuilder) Scan(dest ...interface{}) error { + return b.QueryRow().Scan(dest...) +} + +// SQL methods + +// ToSql builds the query into a SQL string and bound args. +func (b SelectBuilder) ToSql() (string, []interface{}, error) { + data := builder.GetStruct(b).(selectData) + return data.ToSql() +} + +func (b SelectBuilder) MustSql() (string, []interface{}) { + sql, args, err := b.ToSql() + if err != nil { + panic(err) + } + return sql, args +} + +func (b SelectBuilder) toSqlRaw() (string, []interface{}, error) { + data := builder.GetStruct(b).(selectData) + return data.toSqlRaw() +} + +// Prefix adds an expression to the beginning of the query +func (b SelectBuilder) Prefix(sql string, args ...interface{}) SelectBuilder { + return builder.Append(b, "Prefixes", Expr(sql, args...)).(SelectBuilder) +} + +// Distinct adds a DISTINCT clause to the query. +func (b SelectBuilder) Distinct() SelectBuilder { + return b.Options("DISTINCT") +} + +// Options adds select option to the query +func (b SelectBuilder) Options(options ...string) SelectBuilder { + return builder.Extend(b, "Options", options).(SelectBuilder) +} + +// Columns adds result columns to the query. +func (b SelectBuilder) Columns(columns ...string) SelectBuilder { + var parts []interface{} + for _, str := range columns { + parts = append(parts, newPart(str)) + } + return builder.Extend(b, "Columns", parts).(SelectBuilder) +} + +// Column adds a result column to the query. +// Unlike Columns, Column accepts args which will be bound to placeholders in +// the columns string, for example: +// Column("IF(col IN ("+squirrel.Placeholders(3)+"), 1, 0) as col", 1, 2, 3) +func (b SelectBuilder) Column(column interface{}, args ...interface{}) SelectBuilder { + return builder.Append(b, "Columns", newPart(column, args...)).(SelectBuilder) +} + +// From sets the FROM clause of the query. +func (b SelectBuilder) From(from string) SelectBuilder { + return builder.Set(b, "From", newPart(from)).(SelectBuilder) +} + +// FromSelect sets a subquery into the FROM clause of the query. +func (b SelectBuilder) FromSelect(from SelectBuilder, alias string) SelectBuilder { + return builder.Set(b, "From", Alias(from, alias)).(SelectBuilder) +} + +// JoinClause adds a join clause to the query. +func (b SelectBuilder) JoinClause(pred interface{}, args ...interface{}) SelectBuilder { + return builder.Append(b, "Joins", newPart(pred, args...)).(SelectBuilder) +} + +// Join adds a JOIN clause to the query. +func (b SelectBuilder) Join(join string, rest ...interface{}) SelectBuilder { + return b.JoinClause("JOIN "+join, rest...) +} + +// LeftJoin adds a LEFT JOIN clause to the query. +func (b SelectBuilder) LeftJoin(join string, rest ...interface{}) SelectBuilder { + return b.JoinClause("LEFT JOIN "+join, rest...) +} + +// RightJoin adds a RIGHT JOIN clause to the query. +func (b SelectBuilder) RightJoin(join string, rest ...interface{}) SelectBuilder { + return b.JoinClause("RIGHT JOIN "+join, rest...) +} + +// Where adds an expression to the WHERE clause of the query. +// +// Expressions are ANDed together in the generated SQL. +// +// Where accepts several types for its pred argument: +// +// nil OR "" - ignored. +// +// string - SQL expression. +// If the expression has SQL placeholders then a set of arguments must be passed +// as well, one for each placeholder. +// +// map[string]interface{} OR Eq - map of SQL expressions to values. Each key is +// transformed into an expression like " = ?", with the corresponding value +// bound to the placeholder. If the value is nil, the expression will be " +// IS NULL". If the value is an array or slice, the expression will be " IN +// (?,?,...)", with one placeholder for each item in the value. These expressions +// are ANDed together. +// +// Where will panic if pred isn't any of the above types. +func (b SelectBuilder) Where(pred interface{}, args ...interface{}) SelectBuilder { + if pred == nil || pred == "" { + return b + } + return builder.Append(b, "WhereParts", newWherePart(pred, args...)).(SelectBuilder) +} + +// GroupBy adds GROUP BY expressions to the query. +func (b SelectBuilder) GroupBy(groupBys ...string) SelectBuilder { + return builder.Extend(b, "GroupBys", groupBys).(SelectBuilder) +} + +// Having adds an expression to the HAVING clause of the query. +// +// See Where. +func (b SelectBuilder) Having(pred interface{}, rest ...interface{}) SelectBuilder { + return builder.Append(b, "HavingParts", newWherePart(pred, rest...)).(SelectBuilder) +} + +// OrderBy adds ORDER BY expressions to the query. +func (b SelectBuilder) OrderBy(orderBys ...string) SelectBuilder { + return builder.Extend(b, "OrderBys", orderBys).(SelectBuilder) +} + +// Limit sets a LIMIT clause on the query. +func (b SelectBuilder) Limit(limit uint64) SelectBuilder { + return builder.Set(b, "Limit", fmt.Sprintf("%d", limit)).(SelectBuilder) +} + +// Limit ALL allows to access all records with limit +func (b SelectBuilder) RemoveLimit() SelectBuilder { + return builder.Delete(b, "Limit").(SelectBuilder) +} + +// Offset sets a OFFSET clause on the query. +func (b SelectBuilder) Offset(offset uint64) SelectBuilder { + return builder.Set(b, "Offset", fmt.Sprintf("%d", offset)).(SelectBuilder) +} + +// Suffix adds an expression to the end of the query +func (b SelectBuilder) Suffix(sql string, args ...interface{}) SelectBuilder { + return builder.Append(b, "Suffixes", Expr(sql, args...)).(SelectBuilder) +} diff --git a/vendor/github.com/Masterminds/squirrel/select_ctx.go b/vendor/github.com/Masterminds/squirrel/select_ctx.go new file mode 100644 index 0000000000..4c42c13f47 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/select_ctx.go @@ -0,0 +1,69 @@ +// +build go1.8 + +package squirrel + +import ( + "context" + "database/sql" + + "github.com/lann/builder" +) + +func (d *selectData) ExecContext(ctx context.Context) (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(ExecerContext) + if !ok { + return nil, NoContextSupport + } + return ExecContextWith(ctx, ctxRunner, d) +} + +func (d *selectData) QueryContext(ctx context.Context) (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(QueryerContext) + if !ok { + return nil, NoContextSupport + } + return QueryContextWith(ctx, ctxRunner, d) +} + +func (d *selectData) QueryRowContext(ctx context.Context) RowScanner { + if d.RunWith == nil { + return &Row{err: RunnerNotSet} + } + queryRower, ok := d.RunWith.(QueryRowerContext) + if !ok { + if _, ok := d.RunWith.(QueryerContext); !ok { + return &Row{err: RunnerNotQueryRunner} + } + return &Row{err: NoContextSupport} + } + return QueryRowContextWith(ctx, queryRower, d) +} + +// ExecContext builds and ExecContexts the query with the Runner set by RunWith. +func (b SelectBuilder) ExecContext(ctx context.Context) (sql.Result, error) { + data := builder.GetStruct(b).(selectData) + return data.ExecContext(ctx) +} + +// QueryContext builds and QueryContexts the query with the Runner set by RunWith. +func (b SelectBuilder) QueryContext(ctx context.Context) (*sql.Rows, error) { + data := builder.GetStruct(b).(selectData) + return data.QueryContext(ctx) +} + +// QueryRowContext builds and QueryRowContexts the query with the Runner set by RunWith. +func (b SelectBuilder) QueryRowContext(ctx context.Context) RowScanner { + data := builder.GetStruct(b).(selectData) + return data.QueryRowContext(ctx) +} + +// ScanContext is a shortcut for QueryRowContext().Scan. +func (b SelectBuilder) ScanContext(ctx context.Context, dest ...interface{}) error { + return b.QueryRowContext(ctx).Scan(dest...) +} diff --git a/vendor/github.com/Masterminds/squirrel/squirrel.go b/vendor/github.com/Masterminds/squirrel/squirrel.go new file mode 100644 index 0000000000..3190ce7e20 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/squirrel.go @@ -0,0 +1,172 @@ +// Package squirrel provides a fluent SQL generator. +// +// See https://github.com/Masterminds/squirrel for examples. +package squirrel + +import ( + "bytes" + "context" + "database/sql" + "fmt" + "strings" + + "github.com/lann/builder" +) + +// Sqlizer is the interface that wraps the ToSql method. +// +// ToSql returns a SQL representation of the Sqlizer, along with a slice of args +// as passed to e.g. database/sql.Exec. It can also return an error. +type Sqlizer interface { + ToSql() (string, []interface{}, error) +} + +// rawSqlizer is expected to do what Sqlizer does, but without finalizing placeholders. +// This is useful for nested queries. +type rawSqlizer interface { + toSqlRaw() (string, []interface{}, error) +} + +// Execer is the interface that wraps the Exec method. +// +// Exec executes the given query as implemented by database/sql.Exec. +type Execer interface { + Exec(query string, args ...interface{}) (sql.Result, error) +} + +// Queryer is the interface that wraps the Query method. +// +// Query executes the given query as implemented by database/sql.Query. +type Queryer interface { + Query(query string, args ...interface{}) (*sql.Rows, error) +} + +// QueryRower is the interface that wraps the QueryRow method. +// +// QueryRow executes the given query as implemented by database/sql.QueryRow. +type QueryRower interface { + QueryRow(query string, args ...interface{}) RowScanner +} + +// BaseRunner groups the Execer and Queryer interfaces. +type BaseRunner interface { + Execer + Queryer +} + +// Runner groups the Execer, Queryer, and QueryRower interfaces. +type Runner interface { + Execer + Queryer + QueryRower +} + +type stdsql interface { + Query(string, ...interface{}) (*sql.Rows, error) + QueryContext(context.Context, string, ...interface{}) (*sql.Rows, error) + QueryRow(string, ...interface{}) *sql.Row + QueryRowContext(context.Context, string, ...interface{}) *sql.Row + Exec(string, ...interface{}) (sql.Result, error) + ExecContext(context.Context, string, ...interface{}) (sql.Result, error) +} + +type stdsqlRunner struct { + stdsql +} + +func (r *stdsqlRunner) QueryRow(query string, args ...interface{}) RowScanner { + return r.stdsql.QueryRow(query, args...) +} + +func setRunWith(b interface{}, baseRunner BaseRunner) interface{} { + var runner Runner + switch r := baseRunner.(type) { + case Runner: + runner = r + case stdsql: + runner = &stdsqlRunner{r} + + } + return builder.Set(b, "RunWith", runner) +} + +// RunnerNotSet is returned by methods that need a Runner if it isn't set. +var RunnerNotSet = fmt.Errorf("cannot run; no Runner set (RunWith)") + +// RunnerNotQueryRunner is returned by QueryRow if the RunWith value doesn't implement QueryRower. +var RunnerNotQueryRunner = fmt.Errorf("cannot QueryRow; Runner is not a QueryRower") + +// ExecWith Execs the SQL returned by s with db. +func ExecWith(db Execer, s Sqlizer) (res sql.Result, err error) { + query, args, err := s.ToSql() + if err != nil { + return + } + return db.Exec(query, args...) +} + +// QueryWith Querys the SQL returned by s with db. +func QueryWith(db Queryer, s Sqlizer) (rows *sql.Rows, err error) { + query, args, err := s.ToSql() + if err != nil { + return + } + return db.Query(query, args...) +} + +// QueryRowWith QueryRows the SQL returned by s with db. +func QueryRowWith(db QueryRower, s Sqlizer) RowScanner { + query, args, err := s.ToSql() + return &Row{RowScanner: db.QueryRow(query, args...), err: err} +} + +// DebugSqlizer calls ToSql on s and shows the approximate SQL to be executed +// +// If ToSql returns an error, the result of this method will look like: +// "[ToSql error: %s]" or "[DebugSqlizer error: %s]" +// +// IMPORTANT: As its name suggests, this function should only be used for +// debugging. While the string result *might* be valid SQL, this function does +// not try very hard to ensure it. Additionally, executing the output of this +// function with any untrusted user input is certainly insecure. +func DebugSqlizer(s Sqlizer) string { + sql, args, err := s.ToSql() + if err != nil { + return fmt.Sprintf("[ToSql error: %s]", err) + } + + // TODO: dedupe this with placeholder.go + buf := &bytes.Buffer{} + i := 0 + for { + p := strings.Index(sql, "?") + if p == -1 { + break + } + if len(sql[p:]) > 1 && sql[p:p+2] == "??" { // escape ?? => ? + buf.WriteString(sql[:p]) + buf.WriteString("?") + if len(sql[p:]) == 1 { + break + } + sql = sql[p+2:] + } else { + if i+1 > len(args) { + return fmt.Sprintf( + "[DebugSqlizer error: too many placeholders in %#v for %d args]", + sql, len(args)) + } + buf.WriteString(sql[:p]) + fmt.Fprintf(buf, "'%v'", args[i]) + sql = sql[p+1:] + i++ + } + } + if i < len(args) { + return fmt.Sprintf( + "[DebugSqlizer error: not enough placeholders in %#v for %d args]", + sql, len(args)) + } + buf.WriteString(sql) + return buf.String() +} diff --git a/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go b/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go new file mode 100644 index 0000000000..9a50067d9f --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go @@ -0,0 +1,61 @@ +// +build go1.8 + +package squirrel + +import ( + "context" + "database/sql" + "errors" +) + +// NoContextSupport is returned if a db doesn't support Context. +var NoContextSupport = errors.New("DB does not support Context") + +// ExecerContext is the interface that wraps the ExecContext method. +// +// Exec executes the given query as implemented by database/sql.ExecContext. +type ExecerContext interface { + ExecContext(ctx context.Context, query string, args ...interface{}) (sql.Result, error) +} + +// QueryerContext is the interface that wraps the QueryContext method. +// +// QueryContext executes the given query as implemented by database/sql.QueryContext. +type QueryerContext interface { + QueryContext(ctx context.Context, query string, args ...interface{}) (*sql.Rows, error) +} + +// QueryRowerContext is the interface that wraps the QueryRowContext method. +// +// QueryRowContext executes the given query as implemented by database/sql.QueryRowContext. +type QueryRowerContext interface { + QueryRowContext(ctx context.Context, query string, args ...interface{}) RowScanner +} + +func (r *stdsqlRunner) QueryRowContext(ctx context.Context, query string, args ...interface{}) RowScanner { + return r.stdsql.QueryRowContext(ctx, query, args...) +} + +// ExecContextWith ExecContexts the SQL returned by s with db. +func ExecContextWith(ctx context.Context, db ExecerContext, s Sqlizer) (res sql.Result, err error) { + query, args, err := s.ToSql() + if err != nil { + return + } + return db.ExecContext(ctx, query, args...) +} + +// QueryContextWith QueryContexts the SQL returned by s with db. +func QueryContextWith(ctx context.Context, db QueryerContext, s Sqlizer) (rows *sql.Rows, err error) { + query, args, err := s.ToSql() + if err != nil { + return + } + return db.QueryContext(ctx, query, args...) +} + +// QueryRowContextWith QueryRowContexts the SQL returned by s with db. +func QueryRowContextWith(ctx context.Context, db QueryRowerContext, s Sqlizer) RowScanner { + query, args, err := s.ToSql() + return &Row{RowScanner: db.QueryRowContext(ctx, query, args...), err: err} +} diff --git a/vendor/github.com/Masterminds/squirrel/statement.go b/vendor/github.com/Masterminds/squirrel/statement.go new file mode 100644 index 0000000000..275388f630 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/statement.go @@ -0,0 +1,83 @@ +package squirrel + +import "github.com/lann/builder" + +// StatementBuilderType is the type of StatementBuilder. +type StatementBuilderType builder.Builder + +// Select returns a SelectBuilder for this StatementBuilderType. +func (b StatementBuilderType) Select(columns ...string) SelectBuilder { + return SelectBuilder(b).Columns(columns...) +} + +// Insert returns a InsertBuilder for this StatementBuilderType. +func (b StatementBuilderType) Insert(into string) InsertBuilder { + return InsertBuilder(b).Into(into) +} + +// Update returns a UpdateBuilder for this StatementBuilderType. +func (b StatementBuilderType) Update(table string) UpdateBuilder { + return UpdateBuilder(b).Table(table) +} + +// Delete returns a DeleteBuilder for this StatementBuilderType. +func (b StatementBuilderType) Delete(from string) DeleteBuilder { + return DeleteBuilder(b).From(from) +} + +// PlaceholderFormat sets the PlaceholderFormat field for any child builders. +func (b StatementBuilderType) PlaceholderFormat(f PlaceholderFormat) StatementBuilderType { + return builder.Set(b, "PlaceholderFormat", f).(StatementBuilderType) +} + +// RunWith sets the RunWith field for any child builders. +func (b StatementBuilderType) RunWith(runner BaseRunner) StatementBuilderType { + return setRunWith(b, runner).(StatementBuilderType) +} + +// StatementBuilder is a parent builder for other builders, e.g. SelectBuilder. +var StatementBuilder = StatementBuilderType(builder.EmptyBuilder).PlaceholderFormat(Question) + +// Select returns a new SelectBuilder, optionally setting some result columns. +// +// See SelectBuilder.Columns. +func Select(columns ...string) SelectBuilder { + return StatementBuilder.Select(columns...) +} + +// Insert returns a new InsertBuilder with the given table name. +// +// See InsertBuilder.Into. +func Insert(into string) InsertBuilder { + return StatementBuilder.Insert(into) +} + +// Update returns a new UpdateBuilder with the given table name. +// +// See UpdateBuilder.Table. +func Update(table string) UpdateBuilder { + return StatementBuilder.Update(table) +} + +// Delete returns a new DeleteBuilder with the given table name. +// +// See DeleteBuilder.Table. +func Delete(from string) DeleteBuilder { + return StatementBuilder.Delete(from) +} + +// Case returns a new CaseBuilder +// "what" represents case value +func Case(what ...interface{}) CaseBuilder { + b := CaseBuilder(builder.EmptyBuilder) + + switch len(what) { + case 0: + case 1: + b = b.what(what[0]) + default: + b = b.what(newPart(what[0], what[1:]...)) + + } + return b +} diff --git a/vendor/github.com/Masterminds/squirrel/stmtcacher.go b/vendor/github.com/Masterminds/squirrel/stmtcacher.go new file mode 100644 index 0000000000..2540565e59 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/stmtcacher.go @@ -0,0 +1,85 @@ +package squirrel + +import ( + "database/sql" + "sync" +) + +// Prepareer is the interface that wraps the Prepare method. +// +// Prepare executes the given query as implemented by database/sql.Prepare. +type Preparer interface { + Prepare(query string) (*sql.Stmt, error) +} + +// DBProxy groups the Execer, Queryer, QueryRower, and Preparer interfaces. +type DBProxy interface { + Execer + Queryer + QueryRower + Preparer +} + +// NOTE: NewStmtCacher is defined in stmtcacher_ctx.go (Go >= 1.8) or stmtcacher_noctx.go (Go < 1.8). + +type stmtCacher struct { + prep Preparer + cache map[string]*sql.Stmt + mu sync.Mutex +} + +func (sc *stmtCacher) Prepare(query string) (*sql.Stmt, error) { + sc.mu.Lock() + defer sc.mu.Unlock() + stmt, ok := sc.cache[query] + if ok { + return stmt, nil + } + stmt, err := sc.prep.Prepare(query) + if err == nil { + sc.cache[query] = stmt + } + return stmt, err +} + +func (sc *stmtCacher) Exec(query string, args ...interface{}) (res sql.Result, err error) { + stmt, err := sc.Prepare(query) + if err != nil { + return + } + return stmt.Exec(args...) +} + +func (sc *stmtCacher) Query(query string, args ...interface{}) (rows *sql.Rows, err error) { + stmt, err := sc.Prepare(query) + if err != nil { + return + } + return stmt.Query(args...) +} + +func (sc *stmtCacher) QueryRow(query string, args ...interface{}) RowScanner { + stmt, err := sc.Prepare(query) + if err != nil { + return &Row{err: err} + } + return stmt.QueryRow(args...) +} + +type DBProxyBeginner interface { + DBProxy + Begin() (*sql.Tx, error) +} + +type stmtCacheProxy struct { + DBProxy + db *sql.DB +} + +func NewStmtCacheProxy(db *sql.DB) DBProxyBeginner { + return &stmtCacheProxy{DBProxy: NewStmtCacher(db), db: db} +} + +func (sp *stmtCacheProxy) Begin() (*sql.Tx, error) { + return sp.db.Begin() +} diff --git a/vendor/github.com/Masterminds/squirrel/stmtcacher_ctx.go b/vendor/github.com/Masterminds/squirrel/stmtcacher_ctx.go new file mode 100644 index 0000000000..2ad51e4944 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/stmtcacher_ctx.go @@ -0,0 +1,74 @@ +// +build go1.8 + +package squirrel + +import ( + "context" + "database/sql" +) + +// PrepareerContext is the interface that wraps the Prepare and PrepareContext methods. +// +// Prepare executes the given query as implemented by database/sql.Prepare. +// PrepareContext executes the given query as implemented by database/sql.PrepareContext. +type PreparerContext interface { + Preparer + PrepareContext(ctx context.Context, query string) (*sql.Stmt, error) +} + +// DBProxyContext groups the Execer, Queryer, QueryRower and PreparerContext interfaces. +type DBProxyContext interface { + Execer + Queryer + QueryRower + PreparerContext +} + +// NewStmtCacher returns a DBProxy wrapping prep that caches Prepared Stmts. +// +// Stmts are cached based on the string value of their queries. +func NewStmtCacher(prep PreparerContext) DBProxyContext { + return &stmtCacher{prep: prep, cache: make(map[string]*sql.Stmt)} +} + +func (sc *stmtCacher) PrepareContext(ctx context.Context, query string) (*sql.Stmt, error) { + ctxPrep, ok := sc.prep.(PreparerContext) + if !ok { + return nil, NoContextSupport + } + sc.mu.Lock() + defer sc.mu.Unlock() + stmt, ok := sc.cache[query] + if ok { + return stmt, nil + } + stmt, err := ctxPrep.PrepareContext(ctx, query) + if err == nil { + sc.cache[query] = stmt + } + return stmt, err +} + +func (sc *stmtCacher) ExecContext(ctx context.Context, query string, args ...interface{}) (res sql.Result, err error) { + stmt, err := sc.PrepareContext(ctx, query) + if err != nil { + return + } + return stmt.ExecContext(ctx, args...) +} + +func (sc *stmtCacher) QueryContext(ctx context.Context, query string, args ...interface{}) (rows *sql.Rows, err error) { + stmt, err := sc.PrepareContext(ctx, query) + if err != nil { + return + } + return stmt.QueryContext(ctx, args...) +} + +func (sc *stmtCacher) QueryRowContext(ctx context.Context, query string, args ...interface{}) RowScanner { + stmt, err := sc.PrepareContext(ctx, query) + if err != nil { + return &Row{err: err} + } + return stmt.QueryRowContext(ctx, args...) +} diff --git a/vendor/github.com/Masterminds/squirrel/stmtcacher_noctx.go b/vendor/github.com/Masterminds/squirrel/stmtcacher_noctx.go new file mode 100644 index 0000000000..f89f5b2c2d --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/stmtcacher_noctx.go @@ -0,0 +1,14 @@ +// +build !go1.8 + +package squirrel + +import ( + "database/sql" +) + +// NewStmtCacher returns a DBProxy wrapping prep that caches Prepared Stmts. +// +// Stmts are cached based on the string value of their queries. +func NewStmtCacher(prep Preparer) DBProxy { + return &stmtCacher{prep: prep, cache: make(map[string]*sql.Stmt)} +} diff --git a/vendor/github.com/Masterminds/squirrel/update.go b/vendor/github.com/Masterminds/squirrel/update.go new file mode 100644 index 0000000000..682906bc05 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/update.go @@ -0,0 +1,232 @@ +package squirrel + +import ( + "bytes" + "database/sql" + "fmt" + "sort" + "strings" + + "github.com/lann/builder" +) + +type updateData struct { + PlaceholderFormat PlaceholderFormat + RunWith BaseRunner + Prefixes exprs + Table string + SetClauses []setClause + WhereParts []Sqlizer + OrderBys []string + Limit string + Offset string + Suffixes exprs +} + +type setClause struct { + column string + value interface{} +} + +func (d *updateData) Exec() (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return ExecWith(d.RunWith, d) +} + +func (d *updateData) Query() (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + return QueryWith(d.RunWith, d) +} + +func (d *updateData) QueryRow() RowScanner { + if d.RunWith == nil { + return &Row{err: RunnerNotSet} + } + queryRower, ok := d.RunWith.(QueryRower) + if !ok { + return &Row{err: RunnerNotQueryRunner} + } + return QueryRowWith(queryRower, d) +} + +func (d *updateData) ToSql() (sqlStr string, args []interface{}, err error) { + if len(d.Table) == 0 { + err = fmt.Errorf("update statements must specify a table") + return + } + if len(d.SetClauses) == 0 { + err = fmt.Errorf("update statements must have at least one Set clause") + return + } + + sql := &bytes.Buffer{} + + if len(d.Prefixes) > 0 { + args, _ = d.Prefixes.AppendToSql(sql, " ", args) + sql.WriteString(" ") + } + + sql.WriteString("UPDATE ") + sql.WriteString(d.Table) + + sql.WriteString(" SET ") + setSqls := make([]string, len(d.SetClauses)) + for i, setClause := range d.SetClauses { + var valSql string + e, isExpr := setClause.value.(expr) + if isExpr { + valSql = e.sql + args = append(args, e.args...) + } else { + valSql = "?" + args = append(args, setClause.value) + } + setSqls[i] = fmt.Sprintf("%s = %s", setClause.column, valSql) + } + sql.WriteString(strings.Join(setSqls, ", ")) + + if len(d.WhereParts) > 0 { + sql.WriteString(" WHERE ") + args, err = appendToSql(d.WhereParts, sql, " AND ", args) + if err != nil { + return + } + } + + if len(d.OrderBys) > 0 { + sql.WriteString(" ORDER BY ") + sql.WriteString(strings.Join(d.OrderBys, ", ")) + } + + if len(d.Limit) > 0 { + sql.WriteString(" LIMIT ") + sql.WriteString(d.Limit) + } + + if len(d.Offset) > 0 { + sql.WriteString(" OFFSET ") + sql.WriteString(d.Offset) + } + + if len(d.Suffixes) > 0 { + sql.WriteString(" ") + args, _ = d.Suffixes.AppendToSql(sql, " ", args) + } + + sqlStr, err = d.PlaceholderFormat.ReplacePlaceholders(sql.String()) + return +} + +// Builder + +// UpdateBuilder builds SQL UPDATE statements. +type UpdateBuilder builder.Builder + +func init() { + builder.Register(UpdateBuilder{}, updateData{}) +} + +// Format methods + +// PlaceholderFormat sets PlaceholderFormat (e.g. Question or Dollar) for the +// query. +func (b UpdateBuilder) PlaceholderFormat(f PlaceholderFormat) UpdateBuilder { + return builder.Set(b, "PlaceholderFormat", f).(UpdateBuilder) +} + +// Runner methods + +// RunWith sets a Runner (like database/sql.DB) to be used with e.g. Exec. +func (b UpdateBuilder) RunWith(runner BaseRunner) UpdateBuilder { + return setRunWith(b, runner).(UpdateBuilder) +} + +// Exec builds and Execs the query with the Runner set by RunWith. +func (b UpdateBuilder) Exec() (sql.Result, error) { + data := builder.GetStruct(b).(updateData) + return data.Exec() +} + +func (b UpdateBuilder) Query() (*sql.Rows, error) { + data := builder.GetStruct(b).(updateData) + return data.Query() +} + +func (b UpdateBuilder) QueryRow() RowScanner { + data := builder.GetStruct(b).(updateData) + return data.QueryRow() +} + +func (b UpdateBuilder) Scan(dest ...interface{}) error { + return b.QueryRow().Scan(dest...) +} + +// SQL methods + +// ToSql builds the query into a SQL string and bound args. +func (b UpdateBuilder) ToSql() (string, []interface{}, error) { + data := builder.GetStruct(b).(updateData) + return data.ToSql() +} + +// Prefix adds an expression to the beginning of the query +func (b UpdateBuilder) Prefix(sql string, args ...interface{}) UpdateBuilder { + return builder.Append(b, "Prefixes", Expr(sql, args...)).(UpdateBuilder) +} + +// Table sets the table to be updated. +func (b UpdateBuilder) Table(table string) UpdateBuilder { + return builder.Set(b, "Table", table).(UpdateBuilder) +} + +// Set adds SET clauses to the query. +func (b UpdateBuilder) Set(column string, value interface{}) UpdateBuilder { + return builder.Append(b, "SetClauses", setClause{column: column, value: value}).(UpdateBuilder) +} + +// SetMap is a convenience method which calls .Set for each key/value pair in clauses. +func (b UpdateBuilder) SetMap(clauses map[string]interface{}) UpdateBuilder { + keys := make([]string, len(clauses)) + i := 0 + for key := range clauses { + keys[i] = key + i++ + } + sort.Strings(keys) + for _, key := range keys { + val, _ := clauses[key] + b = b.Set(key, val) + } + return b +} + +// Where adds WHERE expressions to the query. +// +// See SelectBuilder.Where for more information. +func (b UpdateBuilder) Where(pred interface{}, args ...interface{}) UpdateBuilder { + return builder.Append(b, "WhereParts", newWherePart(pred, args...)).(UpdateBuilder) +} + +// OrderBy adds ORDER BY expressions to the query. +func (b UpdateBuilder) OrderBy(orderBys ...string) UpdateBuilder { + return builder.Extend(b, "OrderBys", orderBys).(UpdateBuilder) +} + +// Limit sets a LIMIT clause on the query. +func (b UpdateBuilder) Limit(limit uint64) UpdateBuilder { + return builder.Set(b, "Limit", fmt.Sprintf("%d", limit)).(UpdateBuilder) +} + +// Offset sets a OFFSET clause on the query. +func (b UpdateBuilder) Offset(offset uint64) UpdateBuilder { + return builder.Set(b, "Offset", fmt.Sprintf("%d", offset)).(UpdateBuilder) +} + +// Suffix adds an expression to the end of the query +func (b UpdateBuilder) Suffix(sql string, args ...interface{}) UpdateBuilder { + return builder.Append(b, "Suffixes", Expr(sql, args...)).(UpdateBuilder) +} diff --git a/vendor/github.com/Masterminds/squirrel/update_ctx.go b/vendor/github.com/Masterminds/squirrel/update_ctx.go new file mode 100644 index 0000000000..ad479f96f4 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/update_ctx.go @@ -0,0 +1,69 @@ +// +build go1.8 + +package squirrel + +import ( + "context" + "database/sql" + + "github.com/lann/builder" +) + +func (d *updateData) ExecContext(ctx context.Context) (sql.Result, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(ExecerContext) + if !ok { + return nil, NoContextSupport + } + return ExecContextWith(ctx, ctxRunner, d) +} + +func (d *updateData) QueryContext(ctx context.Context) (*sql.Rows, error) { + if d.RunWith == nil { + return nil, RunnerNotSet + } + ctxRunner, ok := d.RunWith.(QueryerContext) + if !ok { + return nil, NoContextSupport + } + return QueryContextWith(ctx, ctxRunner, d) +} + +func (d *updateData) QueryRowContext(ctx context.Context) RowScanner { + if d.RunWith == nil { + return &Row{err: RunnerNotSet} + } + queryRower, ok := d.RunWith.(QueryRowerContext) + if !ok { + if _, ok := d.RunWith.(QueryerContext); !ok { + return &Row{err: RunnerNotQueryRunner} + } + return &Row{err: NoContextSupport} + } + return QueryRowContextWith(ctx, queryRower, d) +} + +// ExecContext builds and ExecContexts the query with the Runner set by RunWith. +func (b UpdateBuilder) ExecContext(ctx context.Context) (sql.Result, error) { + data := builder.GetStruct(b).(updateData) + return data.ExecContext(ctx) +} + +// QueryContext builds and QueryContexts the query with the Runner set by RunWith. +func (b UpdateBuilder) QueryContext(ctx context.Context) (*sql.Rows, error) { + data := builder.GetStruct(b).(updateData) + return data.QueryContext(ctx) +} + +// QueryRowContext builds and QueryRowContexts the query with the Runner set by RunWith. +func (b UpdateBuilder) QueryRowContext(ctx context.Context) RowScanner { + data := builder.GetStruct(b).(updateData) + return data.QueryRowContext(ctx) +} + +// ScanContext is a shortcut for QueryRowContext().Scan. +func (b UpdateBuilder) ScanContext(ctx context.Context, dest ...interface{}) error { + return b.QueryRowContext(ctx).Scan(dest...) +} diff --git a/vendor/github.com/Masterminds/squirrel/where.go b/vendor/github.com/Masterminds/squirrel/where.go new file mode 100644 index 0000000000..976b63ace4 --- /dev/null +++ b/vendor/github.com/Masterminds/squirrel/where.go @@ -0,0 +1,30 @@ +package squirrel + +import ( + "fmt" +) + +type wherePart part + +func newWherePart(pred interface{}, args ...interface{}) Sqlizer { + return &wherePart{pred: pred, args: args} +} + +func (p wherePart) ToSql() (sql string, args []interface{}, err error) { + switch pred := p.pred.(type) { + case nil: + // no-op + case rawSqlizer: + return pred.toSqlRaw() + case Sqlizer: + return pred.ToSql() + case map[string]interface{}: + return Eq(pred).ToSql() + case string: + sql = pred + args = p.args + default: + err = fmt.Errorf("expected string-keyed map or string, not %T", pred) + } + return +} diff --git a/vendor/github.com/lann/builder/.gitignore b/vendor/github.com/lann/builder/.gitignore new file mode 100644 index 0000000000..f54eb28f6e --- /dev/null +++ b/vendor/github.com/lann/builder/.gitignore @@ -0,0 +1,2 @@ +*~ +\#*# diff --git a/vendor/github.com/lann/builder/.travis.yml b/vendor/github.com/lann/builder/.travis.yml new file mode 100644 index 0000000000..c8860f69bc --- /dev/null +++ b/vendor/github.com/lann/builder/.travis.yml @@ -0,0 +1,7 @@ +language: go + +go: + - '1.8' + - '1.9' + - '1.10' + - tip diff --git a/vendor/github.com/lann/builder/LICENSE b/vendor/github.com/lann/builder/LICENSE new file mode 100644 index 0000000000..a109e8051c --- /dev/null +++ b/vendor/github.com/lann/builder/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2014-2015 Lann Martin + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/vendor/github.com/lann/builder/README.md b/vendor/github.com/lann/builder/README.md new file mode 100644 index 0000000000..3b18550bb9 --- /dev/null +++ b/vendor/github.com/lann/builder/README.md @@ -0,0 +1,68 @@ +# Builder - fluent immutable builders for Go + +[![GoDoc](https://godoc.org/github.com/lann/builder?status.png)](https://godoc.org/github.com/lann/builder) +[![Build Status](https://travis-ci.org/lann/builder.png?branch=master)](https://travis-ci.org/lann/builder) + +Builder was originally written for +[Squirrel](https://github.com/lann/squirrel), a fluent SQL generator. It +is probably the best example of Builder in action. + +Builder helps you write **fluent** DSLs for your libraries with method chaining: + +```go +resp := ReqBuilder. + Url("http://golang.org"). + Header("User-Agent", "Builder"). + Get() +``` + +Builder uses **immutable** persistent data structures +([these](https://github.com/mndrix/ps), specifically) +so that each step in your method chain can be reused: + +```go +build := WordBuilder.AddLetters("Build") +builder := build.AddLetters("er") +building := build.AddLetters("ing") +``` + +Builder makes it easy to **build** structs using the **builder** pattern +(*surprise!*): + +```go +import "github.com/lann/builder" + +type Muppet struct { + Name string + Friends []string +} + +type muppetBuilder builder.Builder + +func (b muppetBuilder) Name(name string) muppetBuilder { + return builder.Set(b, "Name", name).(muppetBuilder) +} + +func (b muppetBuilder) AddFriend(friend string) muppetBuilder { + return builder.Append(b, "Friends", friend).(muppetBuilder) +} + +func (b muppetBuilder) Build() Muppet { + return builder.GetStruct(b).(Muppet) +} + +var MuppetBuilder = builder.Register(muppetBuilder{}, Muppet{}).(muppetBuilder) +``` +```go +MuppetBuilder. + Name("Beaker"). + AddFriend("Dr. Honeydew"). + Build() + +=> Muppet{Name:"Beaker", Friends:[]string{"Dr. Honeydew"}} +``` + +## License + +Builder is released under the +[MIT License](http://www.opensource.org/licenses/MIT). diff --git a/vendor/github.com/lann/builder/builder.go b/vendor/github.com/lann/builder/builder.go new file mode 100644 index 0000000000..ff621406a4 --- /dev/null +++ b/vendor/github.com/lann/builder/builder.go @@ -0,0 +1,225 @@ +// Package builder provides a method for writing fluent immutable builders. +package builder + +import ( + "github.com/lann/ps" + "go/ast" + "reflect" +) + +// Builder stores a set of named values. +// +// New types can be declared with underlying type Builder and used with the +// functions in this package. See example. +// +// Instances of Builder should be treated as immutable. It is up to the +// implementor to ensure mutable values set on a Builder are not mutated while +// the Builder is in use. +type Builder struct { + builderMap ps.Map +} + +var ( + EmptyBuilder = Builder{ps.NewMap()} + emptyBuilderValue = reflect.ValueOf(EmptyBuilder) +) + +func getBuilderMap(builder interface{}) ps.Map { + b := convert(builder, Builder{}).(Builder) + + if b.builderMap == nil { + return ps.NewMap() + } + + return b.builderMap +} + +// Set returns a copy of the given builder with a new value set for the given +// name. +// +// Set (and all other functions taking a builder in this package) will panic if +// the given builder's underlying type is not Builder. +func Set(builder interface{}, name string, v interface{}) interface{} { + b := Builder{getBuilderMap(builder).Set(name, v)} + return convert(b, builder) +} + +// Delete returns a copy of the given builder with the given named value unset. +func Delete(builder interface{}, name string) interface{} { + b := Builder{getBuilderMap(builder).Delete(name)} + return convert(b, builder) +} + +// Append returns a copy of the given builder with new value(s) appended to the +// named list. If the value was previously unset or set with Set (even to a e.g. +// slice values), the new value(s) will be appended to an empty list. +func Append(builder interface{}, name string, vs ...interface{}) interface{} { + return Extend(builder, name, vs) +} + +// Extend behaves like Append, except it takes a single slice or array value +// which will be concatenated to the named list. +// +// Unlike a variadic call to Append - which requires a []interface{} value - +// Extend accepts slices or arrays of any type. +// +// Extend will panic if the given value is not a slice, array, or nil. +func Extend(builder interface{}, name string, vs interface{}) interface{} { + if vs == nil { + return builder + } + + maybeList, ok := getBuilderMap(builder).Lookup(name) + + var list ps.List + if ok { + list, ok = maybeList.(ps.List) + } + if !ok { + list = ps.NewList() + } + + forEach(vs, func(v interface{}) { + list = list.Cons(v) + }) + + return Set(builder, name, list) +} + +func listToSlice(list ps.List, arrayType reflect.Type) reflect.Value { + size := list.Size() + slice := reflect.MakeSlice(arrayType, size, size) + for i := size - 1; i >= 0; i-- { + val := reflect.ValueOf(list.Head()) + slice.Index(i).Set(val) + list = list.Tail() + } + return slice +} + +var anyArrayType = reflect.TypeOf([]interface{}{}) + +// Get retrieves a single named value from the given builder. +// If the value has not been set, it returns (nil, false). Otherwise, it will +// return (value, true). +// +// If the named value was last set with Append or Extend, the returned value +// will be a slice. If the given Builder has been registered with Register or +// RegisterType and the given name is an exported field of the registered +// struct, the returned slice will have the same type as that field. Otherwise +// the slice will have type []interface{}. It will panic if the given name is a +// registered struct's exported field and the value set on the Builder is not +// assignable to the field. +func Get(builder interface{}, name string) (interface{}, bool) { + val, ok := getBuilderMap(builder).Lookup(name) + if !ok { + return nil, false + } + + list, isList := val.(ps.List) + if isList { + arrayType := anyArrayType + + if ast.IsExported(name) { + structType := getBuilderStructType(reflect.TypeOf(builder)) + if structType != nil { + field, ok := (*structType).FieldByName(name) + if ok { + arrayType = field.Type + } + } + } + + val = listToSlice(list, arrayType).Interface() + } + + return val, true +} + +// GetMap returns a map[string]interface{} of the values set in the given +// builder. +// +// See notes on Get regarding returned slices. +func GetMap(builder interface{}) map[string]interface{} { + m := getBuilderMap(builder) + structType := getBuilderStructType(reflect.TypeOf(builder)) + + ret := make(map[string]interface{}, m.Size()) + + m.ForEach(func(name string, val ps.Any) { + list, isList := val.(ps.List) + if isList { + arrayType := anyArrayType + + if structType != nil { + field, ok := (*structType).FieldByName(name) + if ok { + arrayType = field.Type + } + } + + val = listToSlice(list, arrayType).Interface() + } + + ret[name] = val + }) + + return ret +} + +// GetStruct builds a new struct from the given registered builder. +// It will return nil if the given builder's type has not been registered with +// Register or RegisterValue. +// +// All values set on the builder with names that start with an uppercase letter +// (i.e. which would be exported if they were identifiers) are assigned to the +// corresponding exported fields of the struct. +// +// GetStruct will panic if any of these "exported" values are not assignable to +// their corresponding struct fields. +func GetStruct(builder interface{}) interface{} { + structVal := newBuilderStruct(reflect.TypeOf(builder)) + if structVal == nil { + return nil + } + return scanStruct(builder, structVal) +} + +// GetStructLike builds a new struct from the given builder with the same type +// as the given struct. +// +// All values set on the builder with names that start with an uppercase letter +// (i.e. which would be exported if they were identifiers) are assigned to the +// corresponding exported fields of the struct. +// +// ScanStruct will panic if any of these "exported" values are not assignable to +// their corresponding struct fields. +func GetStructLike(builder interface{}, strct interface{}) interface{} { + structVal := reflect.New(reflect.TypeOf(strct)).Elem() + return scanStruct(builder, &structVal) +} + +func scanStruct(builder interface{}, structVal *reflect.Value) interface{} { + getBuilderMap(builder).ForEach(func(name string, val ps.Any) { + if ast.IsExported(name) { + field := structVal.FieldByName(name) + + var value reflect.Value + switch v := val.(type) { + case nil: + switch field.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Ptr, reflect.Slice: + value = reflect.Zero(field.Type()) + } + // nil is not valid for this Type; Set will panic + case ps.List: + value = listToSlice(v, field.Type()) + default: + value = reflect.ValueOf(val) + } + field.Set(value) + } + }) + + return structVal.Interface() +} diff --git a/vendor/github.com/lann/builder/reflect.go b/vendor/github.com/lann/builder/reflect.go new file mode 100644 index 0000000000..3b236e7918 --- /dev/null +++ b/vendor/github.com/lann/builder/reflect.go @@ -0,0 +1,24 @@ +package builder + +import "reflect" + +func convert(from interface{}, to interface{}) interface{} { + return reflect. + ValueOf(from). + Convert(reflect.TypeOf(to)). + Interface() +} + +func forEach(s interface{}, f func(interface{})) { + val := reflect.ValueOf(s) + + kind := val.Kind() + if kind != reflect.Slice && kind != reflect.Array { + panic(&reflect.ValueError{Method: "builder.forEach", Kind: kind}) + } + + l := val.Len() + for i := 0; i < l; i++ { + f(val.Index(i).Interface()) + } +} diff --git a/vendor/github.com/lann/builder/registry.go b/vendor/github.com/lann/builder/registry.go new file mode 100644 index 0000000000..612845418e --- /dev/null +++ b/vendor/github.com/lann/builder/registry.go @@ -0,0 +1,59 @@ +package builder + +import ( + "reflect" + "sync" +) + +var ( + registry = make(map[reflect.Type]reflect.Type) + registryMux sync.RWMutex +) + +// RegisterType maps the given builderType to a structType. +// This mapping affects the type of slices returned by Get and is required for +// GetStruct to work. +// +// Returns a Value containing an empty instance of the registered builderType. +// +// RegisterType will panic if builderType's underlying type is not Builder or +// if structType's Kind is not Struct. +func RegisterType(builderType reflect.Type, structType reflect.Type) *reflect.Value { + registryMux.Lock() + defer registryMux.Unlock() + structType.NumField() // Panic if structType is not a struct + registry[builderType] = structType + emptyValue := emptyBuilderValue.Convert(builderType) + return &emptyValue +} + +// Register wraps RegisterType, taking instances instead of Types. +// +// Returns an empty instance of the registered builder type which can be used +// as the initial value for builder expressions. See example. +func Register(builderProto, structProto interface{}) interface{} { + empty := RegisterType( + reflect.TypeOf(builderProto), + reflect.TypeOf(structProto), + ).Interface() + return empty +} + +func getBuilderStructType(builderType reflect.Type) *reflect.Type { + registryMux.RLock() + defer registryMux.RUnlock() + structType, ok := registry[builderType] + if !ok { + return nil + } + return &structType +} + +func newBuilderStruct(builderType reflect.Type) *reflect.Value { + structType := getBuilderStructType(builderType) + if structType == nil { + return nil + } + newStruct := reflect.New(*structType).Elem() + return &newStruct +} diff --git a/vendor/github.com/lann/ps/LICENSE b/vendor/github.com/lann/ps/LICENSE new file mode 100644 index 0000000000..69f9ae8d5a --- /dev/null +++ b/vendor/github.com/lann/ps/LICENSE @@ -0,0 +1,7 @@ +Copyright (c) 2013 Michael Hendricks + +Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. diff --git a/vendor/github.com/lann/ps/README.md b/vendor/github.com/lann/ps/README.md new file mode 100644 index 0000000000..325a5b3540 --- /dev/null +++ b/vendor/github.com/lann/ps/README.md @@ -0,0 +1,10 @@ +**This is a stable fork of https://github.com/mndrix/ps; it will not introduce breaking changes.** + +ps +== + +Persistent data structures for Go. See the [full package documentation](http://godoc.org/github.com/lann/ps) + +Install with + + go get github.com/lann/ps diff --git a/vendor/github.com/lann/ps/list.go b/vendor/github.com/lann/ps/list.go new file mode 100644 index 0000000000..48a1cebacf --- /dev/null +++ b/vendor/github.com/lann/ps/list.go @@ -0,0 +1,93 @@ +package ps + +// List is a persistent list of possibly heterogenous values. +type List interface { + // IsNil returns true if the list is empty + IsNil() bool + + // Cons returns a new list with val as the head + Cons(val Any) List + + // Head returns the first element of the list; + // panics if the list is empty + Head() Any + + // Tail returns a list with all elements except the head; + // panics if the list is empty + Tail() List + + // Size returns the list's length. This takes O(1) time. + Size() int + + // ForEach executes a callback for each value in the list. + ForEach(f func(Any)) + + // Reverse returns a list whose elements are in the opposite order as + // the original list. + Reverse() List +} + +// Immutable (i.e. persistent) list +type list struct { + depth int // the number of nodes after, and including, this one + value Any + tail *list +} + +// An empty list shared by all lists +var nilList = &list{} + +// NewList returns a new, empty list. The result is a singly linked +// list implementation. All lists share an empty tail, so allocating +// empty lists is efficient in time and memory. +func NewList() List { + return nilList +} + +func (self *list) IsNil() bool { + return self == nilList +} + +func (self *list) Size() int { + return self.depth +} + +func (tail *list) Cons(val Any) List { + var xs list + xs.depth = tail.depth + 1 + xs.value = val + xs.tail = tail + return &xs +} + +func (self *list) Head() Any { + if self.IsNil() { + panic("Called Head() on an empty list") + } + + return self.value +} + +func (self *list) Tail() List { + if self.IsNil() { + panic("Called Tail() on an empty list") + } + + return self.tail +} + +// ForEach executes a callback for each value in the list +func (self *list) ForEach(f func(Any)) { + if self.IsNil() { + return + } + f(self.Head()) + self.Tail().ForEach(f) +} + +// Reverse returns a list with elements in opposite order as this list +func (self *list) Reverse() List { + reversed := NewList() + self.ForEach(func(v Any) { reversed = reversed.Cons(v) }) + return reversed +} diff --git a/vendor/github.com/lann/ps/map.go b/vendor/github.com/lann/ps/map.go new file mode 100644 index 0000000000..192daec8f3 --- /dev/null +++ b/vendor/github.com/lann/ps/map.go @@ -0,0 +1,311 @@ +// Fully persistent data structures. A persistent data structure is a data +// structure that always preserves the previous version of itself when +// it is modified. Such data structures are effectively immutable, +// as their operations do not update the structure in-place, but instead +// always yield a new structure. +// +// Persistent +// data structures typically share structure among themselves. This allows +// operations to avoid copying the entire data structure. +package ps + +import ( + "bytes" + "fmt" +) + +// Any is a shorthand for Go's verbose interface{} type. +type Any interface{} + +// A Map associates unique keys (type string) with values (type Any). +type Map interface { + // IsNil returns true if the Map is empty + IsNil() bool + + // Set returns a new map in which key and value are associated. + // If the key didn't exist before, it's created; otherwise, the + // associated value is changed. + // This operation is O(log N) in the number of keys. + Set(key string, value Any) Map + + // Delete returns a new map with the association for key, if any, removed. + // This operation is O(log N) in the number of keys. + Delete(key string) Map + + // Lookup returns the value associated with a key, if any. If the key + // exists, the second return value is true; otherwise, false. + // This operation is O(log N) in the number of keys. + Lookup(key string) (Any, bool) + + // Size returns the number of key value pairs in the map. + // This takes O(1) time. + Size() int + + // ForEach executes a callback on each key value pair in the map. + ForEach(f func(key string, val Any)) + + // Keys returns a slice with all keys in this map. + // This operation is O(N) in the number of keys. + Keys() []string + + String() string +} + +// Immutable (i.e. persistent) associative array +const childCount = 8 +const shiftSize = 3 + +type tree struct { + count int + hash uint64 // hash of the key (used for tree balancing) + key string + value Any + children [childCount]*tree +} + +var nilMap = &tree{} + +// Recursively set nilMap's subtrees to point at itself. +// This eliminates all nil pointers in the map structure. +// All map nodes are created by cloning this structure so +// they avoid the problem too. +func init() { + for i := range nilMap.children { + nilMap.children[i] = nilMap + } +} + +// NewMap allocates a new, persistent map from strings to values of +// any type. +// This is currently implemented as a path-copying binary tree. +func NewMap() Map { + return nilMap +} + +func (self *tree) IsNil() bool { + return self == nilMap +} + +// clone returns an exact duplicate of a tree node +func (self *tree) clone() *tree { + var m tree + m = *self + return &m +} + +// constants for FNV-1a hash algorithm +const ( + offset64 uint64 = 14695981039346656037 + prime64 uint64 = 1099511628211 +) + +// hashKey returns a hash code for a given string +func hashKey(key string) uint64 { + hash := offset64 + for _, codepoint := range key { + hash ^= uint64(codepoint) + hash *= prime64 + } + return hash +} + +// Set returns a new map similar to this one but with key and value +// associated. If the key didn't exist, it's created; otherwise, the +// associated value is changed. +func (self *tree) Set(key string, value Any) Map { + hash := hashKey(key) + return setLowLevel(self, hash, hash, key, value) +} + +func setLowLevel(self *tree, partialHash, hash uint64, key string, value Any) *tree { + if self.IsNil() { // an empty tree is easy + m := self.clone() + m.count = 1 + m.hash = hash + m.key = key + m.value = value + return m + } + + if hash != self.hash { + m := self.clone() + i := partialHash % childCount + m.children[i] = setLowLevel(self.children[i], partialHash>>shiftSize, hash, key, value) + recalculateCount(m) + return m + } + + // replacing a key's previous value + m := self.clone() + m.value = value + return m +} + +// modifies a map by recalculating its key count based on the counts +// of its subtrees +func recalculateCount(m *tree) { + count := 0 + for _, t := range m.children { + count += t.Size() + } + m.count = count + 1 // add one to count ourself +} + +func (m *tree) Delete(key string) Map { + hash := hashKey(key) + newMap, _ := deleteLowLevel(m, hash, hash) + return newMap +} + +func deleteLowLevel(self *tree, partialHash, hash uint64) (*tree, bool) { + // empty trees are easy + if self.IsNil() { + return self, false + } + + if hash != self.hash { + i := partialHash % childCount + child, found := deleteLowLevel(self.children[i], partialHash>>shiftSize, hash) + if !found { + return self, false + } + newMap := self.clone() + newMap.children[i] = child + recalculateCount(newMap) + return newMap, true // ? this wasn't in the original code + } + + // we must delete our own node + if self.isLeaf() { // we have no children + return nilMap, true + } + /* + if self.subtreeCount() == 1 { // only one subtree + for _, t := range self.children { + if t != nilMap { + return t, true + } + } + panic("Tree with 1 subtree actually had no subtrees") + } + */ + + // find a node to replace us + i := -1 + size := -1 + for j, t := range self.children { + if t.Size() > size { + i = j + size = t.Size() + } + } + + // make chosen leaf smaller + replacement, child := self.children[i].deleteLeftmost() + newMap := replacement.clone() + for j := range self.children { + if j == i { + newMap.children[j] = child + } else { + newMap.children[j] = self.children[j] + } + } + recalculateCount(newMap) + return newMap, true +} + +// delete the leftmost node in a tree returning the node that +// was deleted and the tree left over after its deletion +func (m *tree) deleteLeftmost() (*tree, *tree) { + if m.isLeaf() { + return m, nilMap + } + + for i, t := range m.children { + if t != nilMap { + deleted, child := t.deleteLeftmost() + newMap := m.clone() + newMap.children[i] = child + recalculateCount(newMap) + return deleted, newMap + } + } + panic("Tree isn't a leaf but also had no children. How does that happen?") +} + +// isLeaf returns true if this is a leaf node +func (m *tree) isLeaf() bool { + return m.Size() == 1 +} + +// returns the number of child subtrees we have +func (m *tree) subtreeCount() int { + count := 0 + for _, t := range m.children { + if t != nilMap { + count++ + } + } + return count +} + +func (m *tree) Lookup(key string) (Any, bool) { + hash := hashKey(key) + return lookupLowLevel(m, hash, hash) +} + +func lookupLowLevel(self *tree, partialHash, hash uint64) (Any, bool) { + if self.IsNil() { // an empty tree is easy + return nil, false + } + + if hash != self.hash { + i := partialHash % childCount + return lookupLowLevel(self.children[i], partialHash>>shiftSize, hash) + } + + // we found it + return self.value, true +} + +func (m *tree) Size() int { + return m.count +} + +func (m *tree) ForEach(f func(key string, val Any)) { + if m.IsNil() { + return + } + + // ourself + f(m.key, m.value) + + // children + for _, t := range m.children { + if t != nilMap { + t.ForEach(f) + } + } +} + +func (m *tree) Keys() []string { + keys := make([]string, m.Size()) + i := 0 + m.ForEach(func(k string, v Any) { + keys[i] = k + i++ + }) + return keys +} + +// make it easier to display maps for debugging +func (m *tree) String() string { + keys := m.Keys() + buf := bytes.NewBufferString("{") + for _, key := range keys { + val, _ := m.Lookup(key) + fmt.Fprintf(buf, "%s: %s, ", key, val) + } + fmt.Fprintf(buf, "}\n") + return buf.String() +} diff --git a/vendor/github.com/lann/ps/profile.sh b/vendor/github.com/lann/ps/profile.sh new file mode 100755 index 0000000000..f03df05a5a --- /dev/null +++ b/vendor/github.com/lann/ps/profile.sh @@ -0,0 +1,3 @@ +#!/bin/sh +go test -c +./ps.test -test.run=none -test.bench=$2 -test.$1profile=$1.profile