package store import ( "context" "database/sql" "strings" "time" ) type ChannelType string const ( ChannelText ChannelType = "text" ChannelVoice ChannelType = "voice" ChannelCategory ChannelType = "category" ChannelDM ChannelType = "dm" ) type Channel struct { ID uint64 GuildID *uint64 Type ChannelType Name string Description string Position int ParentID *uint64 SlowmodeSeconds int NSFW bool Background string VoiceStatus string UserLimit int CreatedAt time.Time } type ChannelOverride struct { ChannelID uint64 TargetType string // role | user TargetID uint64 Allow uint64 Deny uint64 } type CreateChannelParams struct { ID uint64 GuildID *uint64 Type ChannelType Name string Description string Position int ParentID *uint64 SlowmodeSeconds int NSFW bool UserLimit int } const channelColumns = `id, guild_id, type, name, description, position, parent_id, slowmode_seconds, nsfw, background, voice_status, user_limit, created_at` func (s *Store) CreateChannel(ctx context.Context, params CreateChannelParams) (*Channel, error) { if params.ID == 0 { params.ID = s.NextID() } _, err := s.writer.ExecContext(ctx, ` INSERT INTO channels (id, guild_id, type, name, description, position, parent_id, slowmode_seconds, nsfw, user_limit, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, mustID(params.ID), nullableID(params.GuildID), string(params.Type), params.Name, params.Description, params.Position, nullableID(params.ParentID), params.SlowmodeSeconds, boolToInt(params.NSFW), params.UserLimit, s.Now()) if err != nil { return nil, err } return s.GetChannel(ctx, params.ID) } func (s *Store) GetChannel(ctx context.Context, id uint64) (*Channel, error) { return scanChannel(s.reader.QueryRowContext(ctx, `SELECT `+channelColumns+` FROM channels WHERE id = ?`, int64(id))) } func (s *Store) ListGuildChannels(ctx context.Context, guildID uint64) ([]Channel, error) { rows, err := s.reader.QueryContext(ctx, `SELECT `+channelColumns+` FROM channels WHERE guild_id = ? ORDER BY position, id`, int64(guildID)) if err != nil { return nil, err } return collectChannels(rows) } type UpdateChannelParams struct { Name *string Description *string Position *int ParentID *uint64 ClearParent bool SlowmodeSeconds *int NSFW *bool Background *string VoiceStatus *string UserLimit *int } func (s *Store) UpdateChannel(ctx context.Context, id uint64, params UpdateChannelParams) (*Channel, error) { sets := []string{} args := []any{} if params.Name != nil { sets = append(sets, "name = ?") args = append(args, *params.Name) } if params.Description != nil { sets = append(sets, "description = ?") args = append(args, *params.Description) } if params.Position != nil { sets = append(sets, "position = ?") args = append(args, *params.Position) } if params.ParentID != nil { sets = append(sets, "parent_id = ?") args = append(args, int64(*params.ParentID)) } if params.ClearParent { sets = append(sets, "parent_id = NULL") } if params.SlowmodeSeconds != nil { sets = append(sets, "slowmode_seconds = ?") args = append(args, *params.SlowmodeSeconds) } if params.NSFW != nil { sets = append(sets, "nsfw = ?") args = append(args, boolToInt(*params.NSFW)) } if params.Background != nil { sets = append(sets, "background = ?") args = append(args, nullableString(*params.Background)) } if params.VoiceStatus != nil { sets = append(sets, "voice_status = ?") args = append(args, *params.VoiceStatus) } if params.UserLimit != nil { sets = append(sets, "user_limit = ?") args = append(args, *params.UserLimit) } if len(sets) == 0 { return s.GetChannel(ctx, id) } idValue, err := idToInt(id) if err != nil { return nil, err } args = append(args, idValue) result, err := s.writer.ExecContext(ctx, buildQuery(updateChannelTemplate, strings.Join(sets, ", ")), args...) if err != nil { return nil, err } if affected, err := result.RowsAffected(); err == nil && affected == 0 { return nil, ErrNotFound } return s.GetChannel(ctx, id) } func (s *Store) DeleteChannel(ctx context.Context, id uint64) error { idValue, err := idToInt(id) if err != nil { return err } _, err = s.writer.ExecContext(ctx, `DELETE FROM channels WHERE id = ?`, idValue) return err } // SetChannelOverride создаёт или заменяет оверрайд прав на комнату. func (s *Store) SetChannelOverride(ctx context.Context, override ChannelOverride) error { _, err := s.writer.ExecContext(ctx, ` INSERT INTO channel_overrides (channel_id, target_type, target_id, allow, deny) VALUES (?, ?, ?, ?, ?) ON CONFLICT (channel_id, target_type, target_id) DO UPDATE SET allow = excluded.allow, deny = excluded.deny`, int64(override.ChannelID), override.TargetType, int64(override.TargetID), int64(override.Allow), int64(override.Deny)) return err } func (s *Store) ListChannelOverrides(ctx context.Context, channelID uint64) ([]ChannelOverride, error) { rows, err := s.reader.QueryContext(ctx, ` SELECT channel_id, target_type, target_id, allow, deny FROM channel_overrides WHERE channel_id = ?`, int64(channelID)) if err != nil { return nil, err } defer rows.Close() overrides := make([]ChannelOverride, 0, 4) for rows.Next() { var ( override ChannelOverride allow, deny int64 ) if err := rows.Scan(&override.ChannelID, &override.TargetType, &override.TargetID, &allow, &deny); err != nil { return nil, err } override.Allow = intToID(allow) override.Deny = intToID(deny) overrides = append(overrides, override) } return overrides, rows.Err() } func scanChannel(scanner interface{ Scan(...any) error }) (*Channel, error) { var ( channel Channel guildID sql.NullInt64 parentID sql.NullInt64 nsfw int background sql.NullString createdAt string ) err := scanner.Scan(&channel.ID, &guildID, &channel.Type, &channel.Name, &channel.Description, &channel.Position, &parentID, &channel.SlowmodeSeconds, &nsfw, &background, &channel.VoiceStatus, &channel.UserLimit, &createdAt) if err != nil { return nil, mapError(err) } channel.GuildID = optionalID(guildID) channel.ParentID = optionalID(parentID) channel.NSFW = nsfw == 1 if background.Valid { channel.Background = background.String } channel.CreatedAt = parseTimestamp(createdAt) return &channel, nil } func collectChannels(rows *sql.Rows) ([]Channel, error) { defer rows.Close() channels := make([]Channel, 0, 8) for rows.Next() { channel, err := scanChannel(rows) if err != nil { return nil, err } channels = append(channels, *channel) } return channels, rows.Err() }