diff options
| author | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-03-14 14:43:24 -0700 |
|---|---|---|
| committer | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-03-14 14:43:24 -0700 |
| commit | 2cf9f33ef42188dd35eaf4905cd2f05557aecca1 (patch) | |
| tree | e2cb7db1ccd7a4defae39395b76b144d66aacf2c /internal | |
| parent | 0003a0912b04308fbbdfcb2801d0373e6d4de2f0 (diff) | |
Refactor ban management behind BanMgr interface
Extract ban logic into a BanMgr interface with two implementations:
- BanFile: file-based YAML storage with support for IP, username, and
nickname bans (backwards-compatible with legacy format)
- RedisBanMgr: Redis-backed implementation with permanent and temporary
ban support, using fail-safe deny-on-error behavior
This replaces scattered Redis calls in API handlers, transaction
handlers, and server connection logic with unified interface calls,
removing the Redis dependency from the API server constructor and
enabling ban functionality for both file-only and Redis deployments.
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/mobius/api.go | 118 | ||||
| -rw-r--r-- | internal/mobius/ban.go | 171 | ||||
| -rw-r--r-- | internal/mobius/ban_test.go | 152 | ||||
| -rw-r--r-- | internal/mobius/redis_ban.go | 200 | ||||
| -rw-r--r-- | internal/mobius/redis_ban_test.go | 222 | ||||
| -rw-r--r-- | internal/mobius/transaction_handlers.go | 70 |
6 files changed, 828 insertions, 105 deletions
diff --git a/internal/mobius/api.go b/internal/mobius/api.go index 2bce8f8..adc9072 100644 --- a/internal/mobius/api.go +++ b/internal/mobius/api.go @@ -11,7 +11,6 @@ import ( "strings" "github.com/jhalter/mobius/hotline" - "github.com/redis/go-redis/v9" ) type logResponseWriter struct { @@ -41,7 +40,6 @@ type APIServer struct { logger *slog.Logger mux *http.ServeMux apiKey string - redis *redis.Client } func (srv *APIServer) authMiddleware(next http.Handler) http.Handler { @@ -64,22 +62,14 @@ func (srv *APIServer) logMiddleware(next http.Handler) http.Handler { } // NewAPIServer creates a new APIServer instance with the specified configuration. -// It sets up all API routes and middleware, and optionally connects to Redis for persistent storage. -func NewAPIServer(hlServer *hotline.Server, reloadFunc func(), logger *slog.Logger, apiKey string, redisAddr string, redisPassword string, redisDB int) *APIServer { +// It sets up all API routes and middleware. +func NewAPIServer(hlServer *hotline.Server, reloadFunc func(), logger *slog.Logger, apiKey string) *APIServer { srv := APIServer{ hlServer: hlServer, logger: logger, mux: http.NewServeMux(), apiKey: apiKey, } - if redisAddr != "" { - srv.redis = redis.NewClient(&redis.Options{ - Addr: redisAddr, - Password: redisPassword, - DB: redisDB, - }) - hlServer.Redis = srv.redis - } srv.mux.Handle("/api/v1/online", srv.logMiddleware(srv.authMiddleware(http.HandlerFunc(srv.OnlineHandler)))) srv.mux.Handle("/api/v1/ban", srv.logMiddleware(srv.authMiddleware(http.HandlerFunc(srv.BanHandler)))) @@ -91,11 +81,11 @@ func NewAPIServer(hlServer *hotline.Server, reloadFunc func(), logger *slog.Logg srv.mux.Handle("/api/v1/shutdown", srv.logMiddleware(srv.authMiddleware(http.HandlerFunc(srv.ShutdownHandler)))) srv.mux.Handle("/api/v1/stats", srv.logMiddleware(srv.authMiddleware(http.HandlerFunc(srv.RenderStats)))) - if srv.redis != nil { - if err := srv.redis.Del(context.Background(), "mobius:online").Err(); err != nil { - srv.logger.Warn("Failed to clear mobius:online in Redis", "err", err) + if hlServer.Redis != nil { + if err := hlServer.Redis.Del(context.Background(), hotline.RedisKeyOnline).Err(); err != nil { + srv.logger.Warn("Failed to clear online users in Redis", "err", err) } else { - srv.logger.Info("Cleared mobius:online in Redis on startup") + srv.logger.Info("Cleared online users in Redis on startup") } } @@ -107,8 +97,8 @@ func NewAPIServer(hlServer *hotline.Server, reloadFunc func(), logger *slog.Logg func (srv *APIServer) OnlineHandler(w http.ResponseWriter, r *http.Request) { var users []map[string]string - if srv.redis != nil { - members, err := srv.redis.SMembers(r.Context(), "mobius:online").Result() + if srv.hlServer.Redis != nil { + members, err := srv.hlServer.Redis.SMembers(r.Context(), hotline.RedisKeyOnline).Result() if err == nil { for _, m := range members { parts := strings.SplitN(m, ":", 3) @@ -157,25 +147,30 @@ func (srv *APIServer) BanHandler(w http.ResponseWriter, r *http.Request) { return } - if srv.redis != nil { - if req.Username != "" { - srv.redis.SAdd(r.Context(), "mobius:banned:users", req.Username) + if req.Username != "" { + if err := srv.hlServer.BanList.BanUsername(req.Username); err != nil { + http.Error(w, "failed to ban username", http.StatusInternalServerError) + return } - if req.Nickname != "" { - srv.redis.SAdd(r.Context(), "mobius:banned:nicknames", req.Nickname) + } + if req.Nickname != "" { + if err := srv.hlServer.BanList.BanNickname(req.Nickname); err != nil { + http.Error(w, "failed to ban nickname", http.StatusInternalServerError) + return } - if req.IP != "" { - srv.redis.SAdd(r.Context(), "mobius:banned:ips", req.IP) + } + if req.IP != "" { + if err := srv.hlServer.BanList.Add(req.IP, nil); err != nil { + http.Error(w, "failed to ban IP", http.StatusInternalServerError) + return } - } else { - // TODO: Fallback } // Disconnect user if online for _, c := range srv.hlServer.ClientMgr.List() { - if (req.Username != "" && string(c.Account.Login) == req.Username) || + if (req.Username != "" && c.Account.Login == req.Username) || (req.Nickname != "" && string(c.UserName) == req.Nickname) || - (req.IP != "" && c.RemoteAddr == req.IP) { + (req.IP != "" && c.IP() == req.IP) { c.Disconnect() } } @@ -197,18 +192,23 @@ func (srv *APIServer) UnbanHandler(w http.ResponseWriter, r *http.Request) { return } - if srv.redis != nil { - if req.Username != "" { - srv.redis.SRem(r.Context(), "mobius:banned:users", req.Username) + if req.Username != "" { + if err := srv.hlServer.BanList.UnbanUsername(req.Username); err != nil { + http.Error(w, "failed to unban username", http.StatusInternalServerError) + return } - if req.Nickname != "" { - srv.redis.SRem(r.Context(), "mobius:banned:nicknames", req.Nickname) + } + if req.Nickname != "" { + if err := srv.hlServer.BanList.UnbanNickname(req.Nickname); err != nil { + http.Error(w, "failed to unban nickname", http.StatusInternalServerError) + return } - if req.IP != "" { - srv.redis.SRem(r.Context(), "mobius:banned:ips", req.IP) + } + if req.IP != "" { + if err := srv.hlServer.BanList.UnbanIP(req.IP); err != nil { + http.Error(w, "failed to unban IP", http.StatusInternalServerError) + return } - } else { - // TODO: Fallback } _, _ = w.Write([]byte(`{"msg":"unbanned"}`)) @@ -217,46 +217,34 @@ func (srv *APIServer) UnbanHandler(w http.ResponseWriter, r *http.Request) { // ListBannedIPsHandler returns a list of all banned IP addresses. // GET /api/v1/banned/ips func (srv *APIServer) ListBannedIPsHandler(w http.ResponseWriter, r *http.Request) { - if srv.redis != nil { - ips, err := srv.redis.SMembers(r.Context(), "mobius:banned:ips").Result() - if err != nil { - http.Error(w, "failed to fetch banned IPs", http.StatusInternalServerError) - return - } - _ = json.NewEncoder(w).Encode(ips) - } else { - // TODO: Fallback + ips, err := srv.hlServer.BanList.ListBannedIPs() + if err != nil { + http.Error(w, "failed to fetch banned IPs", http.StatusInternalServerError) + return } + _ = json.NewEncoder(w).Encode(ips) } // ListBannedUsernamesHandler returns a list of all banned usernames. // GET /api/v1/banned/usernames func (srv *APIServer) ListBannedUsernamesHandler(w http.ResponseWriter, r *http.Request) { - if srv.redis != nil { - users, err := srv.redis.SMembers(r.Context(), "mobius:banned:users").Result() - if err != nil { - http.Error(w, "failed to fetch banned usernames", http.StatusInternalServerError) - return - } - _ = json.NewEncoder(w).Encode(users) - } else { - // TODO: Fallback + users, err := srv.hlServer.BanList.ListBannedUsernames() + if err != nil { + http.Error(w, "failed to fetch banned usernames", http.StatusInternalServerError) + return } + _ = json.NewEncoder(w).Encode(users) } // ListBannedNicknamesHandler returns a list of all banned nicknames. // GET /api/v1/banned/nicknames func (srv *APIServer) ListBannedNicknamesHandler(w http.ResponseWriter, r *http.Request) { - if srv.redis != nil { - nicks, err := srv.redis.SMembers(r.Context(), "mobius:banned:nicknames").Result() - if err != nil { - http.Error(w, "failed to fetch banned nicknames", http.StatusInternalServerError) - return - } - _ = json.NewEncoder(w).Encode(nicks) - } else { - // TODO: Fallback + nicks, err := srv.hlServer.BanList.ListBannedNicknames() + if err != nil { + http.Error(w, "failed to fetch banned nicknames", http.StatusInternalServerError) + return } + _ = json.NewEncoder(w).Encode(nicks) } // ShutdownHandler gracefully shuts down the server with a custom message. diff --git a/internal/mobius/ban.go b/internal/mobius/ban.go index dcb592e..478dbb0 100644 --- a/internal/mobius/ban.go +++ b/internal/mobius/ban.go @@ -2,25 +2,31 @@ package mobius import ( "fmt" + "io" "os" "path" "sync" "time" + "github.com/jhalter/mobius/hotline" "gopkg.in/yaml.v3" ) type BanFile struct { - banList map[string]*time.Time - filePath string + banList map[string]*time.Time + bannedUsers map[string]bool + bannedNicks map[string]bool + filePath string sync.Mutex } func NewBanFile(path string) (*BanFile, error) { bf := &BanFile{ - filePath: path, - banList: make(map[string]*time.Time), + filePath: path, + banList: make(map[string]*time.Time), + bannedUsers: make(map[string]bool), + bannedNicks: make(map[string]bool), } err := bf.Load() @@ -31,11 +37,19 @@ func NewBanFile(path string) (*BanFile, error) { return bf, nil } +type BanFileData struct { + BanList map[string]*time.Time `yaml:"banList"` + BannedUsers map[string]bool `yaml:"bannedUsers"` + BannedNicks map[string]bool `yaml:"bannedNicks"` +} + func (bf *BanFile) Load() error { bf.Lock() defer bf.Unlock() bf.banList = make(map[string]*time.Time) + bf.bannedUsers = make(map[string]bool) + bf.bannedNicks = make(map[string]bool) fh, err := os.Open(bf.filePath) if os.IsNotExist(err) { @@ -46,21 +60,64 @@ func (bf *BanFile) Load() error { } defer func() { _ = fh.Close() }() - err = yaml.NewDecoder(fh).Decode(&bf.banList) + // Read all file content for proper format detection + content, err := io.ReadAll(fh) + if err != nil { + return fmt.Errorf("read file: %v", err) + } + + // Try to decode as new format first + var data BanFileData + err = yaml.Unmarshal(content, &data) + if err == nil && (data.BanList != nil || data.BannedUsers != nil || data.BannedNicks != nil) { + // Successfully decoded as new format and has actual data + if data.BanList != nil { + bf.banList = data.BanList + } + if data.BannedUsers != nil { + bf.bannedUsers = data.BannedUsers + } + if data.BannedNicks != nil { + bf.bannedNicks = data.BannedNicks + } + return nil + } + + // Try to decode as legacy format (simple map) + var legacyData map[string]*time.Time + err = yaml.Unmarshal(content, &legacyData) if err != nil { return fmt.Errorf("decode yaml: %v", err) } + bf.banList = legacyData + return nil } +// add is the internal implementation that assumes the caller holds the lock. +func (bf *BanFile) add(ip string, until *time.Time) error { + bf.banList[ip] = until + return bf.save() +} + func (bf *BanFile) Add(ip string, until *time.Time) error { bf.Lock() defer bf.Unlock() - bf.banList[ip] = until + return bf.add(ip, until) +} - out, err := yaml.Marshal(bf.banList) +// save persists the ban data to disk. +// Caller must hold bf.Lock(). +func (bf *BanFile) save() error { + data := BanFileData{ + BanList: bf.banList, + BannedUsers: bf.bannedUsers, + BannedNicks: bf.bannedNicks, + } + + out, err := yaml.Marshal(data) if err != nil { return fmt.Errorf("marshal yaml: %v", err) } @@ -83,3 +140,103 @@ func (bf *BanFile) IsBanned(ip string) (bool, *time.Time) { return false, nil } + +// UnbanIP removes an IP from the banned IPs list +func (bf *BanFile) UnbanIP(ip string) error { + bf.Lock() + defer bf.Unlock() + + delete(bf.banList, ip) + return bf.save() +} + +// BanUsername adds a username to the banned users set +func (bf *BanFile) BanUsername(username string) error { + bf.Lock() + defer bf.Unlock() + + bf.bannedUsers[username] = true + return bf.save() +} + +// UnbanUsername removes a username from the banned users set +func (bf *BanFile) UnbanUsername(username string) error { + bf.Lock() + defer bf.Unlock() + + delete(bf.bannedUsers, username) + return bf.save() +} + +// IsUsernameBanned checks if a username is banned +func (bf *BanFile) IsUsernameBanned(username string) bool { + bf.Lock() + defer bf.Unlock() + + return bf.bannedUsers[username] +} + +// BanNickname adds a nickname to the banned nicknames set +func (bf *BanFile) BanNickname(nickname string) error { + bf.Lock() + defer bf.Unlock() + + bf.bannedNicks[nickname] = true + return bf.save() +} + +// UnbanNickname removes a nickname from the banned nicknames set +func (bf *BanFile) UnbanNickname(nickname string) error { + bf.Lock() + defer bf.Unlock() + + delete(bf.bannedNicks, nickname) + return bf.save() +} + +// IsNicknameBanned checks if a nickname is banned +func (bf *BanFile) IsNicknameBanned(nickname string) bool { + bf.Lock() + defer bf.Unlock() + + return bf.bannedNicks[nickname] +} + +// ListBannedIPs returns all banned IP addresses +func (bf *BanFile) ListBannedIPs() ([]string, error) { + bf.Lock() + defer bf.Unlock() + + var ips []string + for ip := range bf.banList { + ips = append(ips, ip) + } + return ips, nil +} + +// ListBannedUsernames returns all banned usernames +func (bf *BanFile) ListBannedUsernames() ([]string, error) { + bf.Lock() + defer bf.Unlock() + + var usernames []string + for username := range bf.bannedUsers { + usernames = append(usernames, username) + } + return usernames, nil +} + +// ListBannedNicknames returns all banned nicknames +func (bf *BanFile) ListBannedNicknames() ([]string, error) { + bf.Lock() + defer bf.Unlock() + + var nicknames []string + for nickname := range bf.bannedNicks { + nicknames = append(nicknames, nickname) + } + return nicknames, nil +} + +// Ensure BanFile implements the BanMgr interface +var _ hotline.BanMgr = (*BanFile)(nil) diff --git a/internal/mobius/ban_test.go b/internal/mobius/ban_test.go index beef715..877fb62 100644 --- a/internal/mobius/ban_test.go +++ b/internal/mobius/ban_test.go @@ -4,11 +4,13 @@ import ( "fmt" "os" "path" + "sort" "sync" "testing" "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestNewBanFile(t *testing.T) { @@ -29,8 +31,10 @@ func TestNewBanFile(t *testing.T) { name: "Valid path with valid content", args: args{path: path.Join(cwd, "test", "config", "Banlist.yaml")}, want: &BanFile{ - filePath: path.Join(cwd, "test", "config", "Banlist.yaml"), - banList: map[string]*time.Time{"192.168.86.29": &testTime}, + filePath: path.Join(cwd, "test", "config", "Banlist.yaml"), + banList: map[string]*time.Time{"192.168.86.29": &testTime}, + bannedUsers: make(map[string]bool), + bannedNicks: make(map[string]bool), }, wantErr: assert.NoError, }, @@ -152,3 +156,147 @@ func TestBanFile_IsBanned(t *testing.T) { }) } } + +func newTempBanFile(t *testing.T) *BanFile { + t.Helper() + tmpDir := t.TempDir() + return &BanFile{ + filePath: path.Join(tmpDir, "banfile.yaml"), + banList: make(map[string]*time.Time), + bannedUsers: make(map[string]bool), + bannedNicks: make(map[string]bool), + } +} + +func TestBanFile_UsernameBanning(t *testing.T) { + bf := newTempBanFile(t) + + // Ban a username. + require.NoError(t, bf.BanUsername("baduser")) + assert.True(t, bf.IsUsernameBanned("baduser")) + assert.False(t, bf.IsUsernameBanned("gooduser")) + + // Persist and reload. + bf2 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf2.Load()) + assert.True(t, bf2.IsUsernameBanned("baduser")) + + // Unban. + require.NoError(t, bf2.UnbanUsername("baduser")) + assert.False(t, bf2.IsUsernameBanned("baduser")) + + // Verify unban persists. + bf3 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf3.Load()) + assert.False(t, bf3.IsUsernameBanned("baduser")) +} + +func TestBanFile_NicknameBanning(t *testing.T) { + bf := newTempBanFile(t) + + // Ban a nickname. + require.NoError(t, bf.BanNickname("troll")) + assert.True(t, bf.IsNicknameBanned("troll")) + assert.False(t, bf.IsNicknameBanned("friend")) + + // Persist and reload. + bf2 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf2.Load()) + assert.True(t, bf2.IsNicknameBanned("troll")) + + // Unban. + require.NoError(t, bf2.UnbanNickname("troll")) + assert.False(t, bf2.IsNicknameBanned("troll")) + + // Verify unban persists. + bf3 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf3.Load()) + assert.False(t, bf3.IsNicknameBanned("troll")) +} + +func TestBanFile_UnbanIP(t *testing.T) { + bf := newTempBanFile(t) + + require.NoError(t, bf.Add("10.0.0.1", nil)) + banned, _ := bf.IsBanned("10.0.0.1") + assert.True(t, banned) + + require.NoError(t, bf.UnbanIP("10.0.0.1")) + banned, _ = bf.IsBanned("10.0.0.1") + assert.False(t, banned) + + // Verify unban persists. + bf2 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf2.Load()) + banned, _ = bf2.IsBanned("10.0.0.1") + assert.False(t, banned) +} + +func TestBanFile_ListOperations(t *testing.T) { + bf := newTempBanFile(t) + + require.NoError(t, bf.Add("1.2.3.4", nil)) + require.NoError(t, bf.Add("5.6.7.8", nil)) + require.NoError(t, bf.BanUsername("user1")) + require.NoError(t, bf.BanUsername("user2")) + require.NoError(t, bf.BanNickname("nick1")) + + ips, err := bf.ListBannedIPs() + require.NoError(t, err) + sort.Strings(ips) + assert.Equal(t, []string{"1.2.3.4", "5.6.7.8"}, ips) + + usernames, err := bf.ListBannedUsernames() + require.NoError(t, err) + sort.Strings(usernames) + assert.Equal(t, []string{"user1", "user2"}, usernames) + + nicknames, err := bf.ListBannedNicknames() + require.NoError(t, err) + assert.Equal(t, []string{"nick1"}, nicknames) +} + +func TestBanFile_PermanentBanViaAdd(t *testing.T) { + bf := newTempBanFile(t) + + require.NoError(t, bf.Add("172.16.0.1", nil)) + + banned, until := bf.IsBanned("172.16.0.1") + assert.True(t, banned) + assert.Nil(t, until, "Add with nil should create a permanent ban") + + // Verify persistence. + bf2 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf2.Load()) + banned, until = bf2.IsBanned("172.16.0.1") + assert.True(t, banned) + assert.Nil(t, until) +} + +func TestBanFile_NewFormatPersistence(t *testing.T) { + bf := newTempBanFile(t) + + expiry := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + require.NoError(t, bf.Add("10.0.0.1", nil)) + require.NoError(t, bf.Add("10.0.0.2", &expiry)) + require.NoError(t, bf.BanUsername("admin")) + require.NoError(t, bf.BanNickname("spammer")) + + // Reload into a fresh BanFile. + bf2 := &BanFile{filePath: bf.filePath} + require.NoError(t, bf2.Load()) + + // Verify IPs. + banned, until := bf2.IsBanned("10.0.0.1") + assert.True(t, banned) + assert.Nil(t, until) + + banned, until = bf2.IsBanned("10.0.0.2") + assert.True(t, banned) + require.NotNil(t, until) + assert.True(t, expiry.Equal(*until)) + + // Verify username and nickname. + assert.True(t, bf2.IsUsernameBanned("admin")) + assert.True(t, bf2.IsNicknameBanned("spammer")) +} diff --git a/internal/mobius/redis_ban.go b/internal/mobius/redis_ban.go new file mode 100644 index 0000000..3c810a4 --- /dev/null +++ b/internal/mobius/redis_ban.go @@ -0,0 +1,200 @@ +package mobius + +import ( + "context" + "log/slog" + "strings" + "time" + + "github.com/jhalter/mobius/hotline" + "github.com/redis/go-redis/v9" +) + +// RedisBanMgr implements the BanMgr interface using Redis as the backend +type RedisBanMgr struct { + client *redis.Client + logger *slog.Logger +} + +// NewRedisBanMgr creates a new Redis-based ban manager +func NewRedisBanMgr(client *redis.Client, logger *slog.Logger) *RedisBanMgr { + return &RedisBanMgr{ + client: client, + logger: logger, + } +} + +// Add adds an IP ban (maintains compatibility with existing interface) +func (r *RedisBanMgr) Add(ip string, until *time.Time) error { + ctx := context.Background() + + if until == nil { + // Permanent ban - add to the permanent ban set + return r.client.SAdd(ctx, hotline.RedisKeyBannedIPs, ip).Err() + } + + // Temporary ban - use a separate key with expiration + tempKey := hotline.RedisKeyTempBannedIPs + ip + duration := time.Until(*until) + if duration <= 0 { + // Already expired, don't add the ban + return nil + } + + return r.client.Set(ctx, tempKey, "1", duration).Err() +} + +// IsBanned checks if an IP is banned (maintains compatibility with existing interface) +// On Redis errors, logs the error and returns true (fail-safe: deny access) +func (r *RedisBanMgr) IsBanned(ip string) (bool, *time.Time) { + ctx := context.Background() + + // Check permanent ban first + banned, err := r.client.SIsMember(ctx, hotline.RedisKeyBannedIPs, ip).Result() + if err != nil { + r.logger.Error("Redis error checking IP ban, denying access (fail-safe)", "ip", ip, "err", err) + return true, nil // Fail-safe: deny access on error + } + if banned { + return true, nil // Permanent ban + } + + // Check temporary ban + tempKey := hotline.RedisKeyTempBannedIPs + ip + exists, err := r.client.Exists(ctx, tempKey).Result() + if err != nil { + r.logger.Error("Redis error checking temp IP ban, denying access (fail-safe)", "ip", ip, "err", err) + return true, nil // Fail-safe: deny access on error + } + if exists > 0 { + // Get TTL to calculate expiration time + ttl, err := r.client.TTL(ctx, tempKey).Result() + if err != nil { + r.logger.Error("Redis error getting TTL for temp ban, denying access (fail-safe)", "ip", ip, "err", err) + return true, nil // Fail-safe: deny access on error + } + if ttl > 0 { + expiration := time.Now().Add(ttl) + return true, &expiration + } + } + + return false, nil +} + +// UnbanIP removes an IP from both permanent and temporary ban lists +func (r *RedisBanMgr) UnbanIP(ip string) error { + ctx := context.Background() + + pipe := r.client.Pipeline() + pipe.SRem(ctx, hotline.RedisKeyBannedIPs, ip) + pipe.Del(ctx, hotline.RedisKeyTempBannedIPs+ip) + + _, err := pipe.Exec(ctx) + return err +} + +// BanUsername adds a username to the banned users set +func (r *RedisBanMgr) BanUsername(username string) error { + ctx := context.Background() + return r.client.SAdd(ctx, hotline.RedisKeyBannedUsers, username).Err() +} + +// UnbanUsername removes a username from the banned users set +func (r *RedisBanMgr) UnbanUsername(username string) error { + ctx := context.Background() + return r.client.SRem(ctx, hotline.RedisKeyBannedUsers, username).Err() +} + +// IsUsernameBanned checks if a username is banned +// On Redis errors, logs the error and returns true (fail-safe: deny access) +func (r *RedisBanMgr) IsUsernameBanned(username string) bool { + ctx := context.Background() + banned, err := r.client.SIsMember(ctx, hotline.RedisKeyBannedUsers, username).Result() + if err != nil { + r.logger.Error("Redis error checking username ban, denying access (fail-safe)", "username", username, "err", err) + return true // Fail-safe: deny access on error + } + return banned +} + +// BanNickname adds a nickname to the banned nicknames set +func (r *RedisBanMgr) BanNickname(nickname string) error { + ctx := context.Background() + return r.client.SAdd(ctx, hotline.RedisKeyBannedNicknames, nickname).Err() +} + +// UnbanNickname removes a nickname from the banned nicknames set +func (r *RedisBanMgr) UnbanNickname(nickname string) error { + ctx := context.Background() + return r.client.SRem(ctx, hotline.RedisKeyBannedNicknames, nickname).Err() +} + +// IsNicknameBanned checks if a nickname is banned +// On Redis errors, logs the error and returns true (fail-safe: deny access) +func (r *RedisBanMgr) IsNicknameBanned(nickname string) bool { + ctx := context.Background() + banned, err := r.client.SIsMember(ctx, hotline.RedisKeyBannedNicknames, nickname).Result() + if err != nil { + r.logger.Error("Redis error checking nickname ban, denying access (fail-safe)", "nickname", nickname, "err", err) + return true // Fail-safe: deny access on error + } + return banned +} + +// ListBannedIPs returns all banned IP addresses (both permanent and temporary) +func (r *RedisBanMgr) ListBannedIPs() ([]string, error) { + ctx := context.Background() + + // Get permanent bans + permanentIPs, err := r.client.SMembers(ctx, hotline.RedisKeyBannedIPs).Result() + if err != nil { + return nil, err + } + + // Get temporary bans by scanning for temp ban keys (non-blocking unlike Keys()) + var tempIPs []string + var cursor uint64 + for { + var keys []string + keys, cursor, err = r.client.Scan(ctx, cursor, hotline.RedisKeyTempBannedIPs+"*", 100).Result() + if err != nil { + return nil, err + } + for _, key := range keys { + ip := strings.TrimPrefix(key, hotline.RedisKeyTempBannedIPs) + tempIPs = append(tempIPs, ip) + } + if cursor == 0 { + break + } + } + + // Combine and deduplicate + allIPs := append(permanentIPs, tempIPs...) + seen := make(map[string]bool) + var uniqueIPs []string + for _, ip := range allIPs { + if !seen[ip] { + seen[ip] = true + uniqueIPs = append(uniqueIPs, ip) + } + } + + return uniqueIPs, nil +} + +// ListBannedUsernames returns all banned usernames +func (r *RedisBanMgr) ListBannedUsernames() ([]string, error) { + ctx := context.Background() + return r.client.SMembers(ctx, hotline.RedisKeyBannedUsers).Result() +} + +// ListBannedNicknames returns all banned nicknames +func (r *RedisBanMgr) ListBannedNicknames() ([]string, error) { + ctx := context.Background() + return r.client.SMembers(ctx, hotline.RedisKeyBannedNicknames).Result() +} + +// Ensure RedisBanMgr implements the BanMgr interface +var _ hotline.BanMgr = (*RedisBanMgr)(nil) diff --git a/internal/mobius/redis_ban_test.go b/internal/mobius/redis_ban_test.go new file mode 100644 index 0000000..6cd3078 --- /dev/null +++ b/internal/mobius/redis_ban_test.go @@ -0,0 +1,222 @@ +package mobius + +import ( + "fmt" + "io" + "log/slog" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestRedisBanMgr_TemporalBans(t *testing.T) { + // Start mini redis server for testing + s, err := miniredis.Run() + require.NoError(t, err) + defer s.Close() + + // Create Redis client + client := redis.NewClient(&redis.Options{ + Addr: s.Addr(), + }) + defer func() { _ = client.Close() }() + + // Create a silent logger for tests + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + banMgr := NewRedisBanMgr(client, logger) + + tests := []struct { + name string + ip string + until *time.Time + expectedBanned bool + expectedUntil *time.Time + setup func() error + cleanup func() error + }{ + { + name: "Permanent ban via Add method", + ip: "192.168.1.100", + until: nil, + expectedBanned: true, + expectedUntil: nil, + }, + { + name: "Temporary ban via Add method", + ip: "192.168.1.101", + until: func() *time.Time { t := time.Now().Add(1 * time.Hour); return &t }(), + expectedBanned: true, + expectedUntil: func() *time.Time { t := time.Now().Add(1 * time.Hour); return &t }(), + }, + { + name: "Expired temporary ban", + ip: "192.168.1.102", + until: func() *time.Time { t := time.Now().Add(-1 * time.Hour); return &t }(), + expectedBanned: false, + expectedUntil: nil, + }, + { + name: "UnbanIP removes both permanent and temporary bans", + ip: "192.168.1.104", + until: nil, // Not used since we have custom setup + expectedBanned: false, + expectedUntil: nil, + setup: func() error { + // Add both permanent and temporary ban + if err := banMgr.Add("192.168.1.104", nil); err != nil { + return err + } + expiration := time.Now().Add(1 * time.Hour) + if err := banMgr.Add("192.168.1.104", &expiration); err != nil { + return err + } + // First verify the IP is banned (permanent takes precedence) + banned, until := banMgr.IsBanned("192.168.1.104") + if !banned || until != nil { + return fmt.Errorf("setup failed: IP should be permanently banned") + } + // Now unban to test the cleanup + return banMgr.UnbanIP("192.168.1.104") + }, + cleanup: func() error { + // Additional cleanup in case something went wrong + return banMgr.UnbanIP("192.168.1.104") + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Setup if needed + if tt.setup != nil { + err := tt.setup() + assert.NoError(t, err) + } else { + // Default setup: add the ban + err := banMgr.Add(tt.ip, tt.until) + assert.NoError(t, err) + } + + // Check ban status + banned, until := banMgr.IsBanned(tt.ip) + assert.Equal(t, tt.expectedBanned, banned) + + if tt.expectedUntil != nil { + assert.NotNil(t, until) + assert.WithinDuration(t, *tt.expectedUntil, *until, 1*time.Second) + } else { + assert.Equal(t, tt.expectedUntil, until) + } + + // Cleanup + if tt.cleanup != nil { + err := tt.cleanup() + assert.NoError(t, err) + } else { + err := banMgr.UnbanIP(tt.ip) + assert.NoError(t, err) + } + }) + } + + t.Run("ListBannedIPs includes both permanent and temporary", func(t *testing.T) { + // Add permanent ban + permanentIP := "192.168.1.105" + err := banMgr.Add(permanentIP, nil) + assert.NoError(t, err) + + // Add temporary ban + temporaryIP := "192.168.1.106" + expiration := time.Now().Add(1 * time.Hour) + err = banMgr.Add(temporaryIP, &expiration) + assert.NoError(t, err) + + // List should include both + ips, err := banMgr.ListBannedIPs() + assert.NoError(t, err) + assert.Contains(t, ips, permanentIP) + assert.Contains(t, ips, temporaryIP) + + // Cleanup + err = banMgr.UnbanIP(permanentIP) + assert.NoError(t, err) + err = banMgr.UnbanIP(temporaryIP) + assert.NoError(t, err) + }) +} + +func TestRedisBanMgr_UserAndNicknameBans(t *testing.T) { + // Start mini redis server for testing + s, err := miniredis.Run() + require.NoError(t, err) + defer s.Close() + + // Create Redis client + client := redis.NewClient(&redis.Options{ + Addr: s.Addr(), + }) + defer func() { _ = client.Close() }() + + // Create a silent logger for tests + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + banMgr := NewRedisBanMgr(client, logger) + + tests := []struct { + name string + banType string + value string + banFunc func(string) error + unbanFunc func(string) error + isBannedFunc func(string) bool + listFunc func() ([]string, error) + }{ + { + name: "User banning", + banType: "username", + value: "testuser", + banFunc: banMgr.BanUsername, + unbanFunc: banMgr.UnbanUsername, + isBannedFunc: banMgr.IsUsernameBanned, + listFunc: banMgr.ListBannedUsernames, + }, + { + name: "Nickname banning", + banType: "nickname", + value: "testnick", + banFunc: banMgr.BanNickname, + unbanFunc: banMgr.UnbanNickname, + isBannedFunc: banMgr.IsNicknameBanned, + listFunc: banMgr.ListBannedNicknames, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Initially not banned + assert.False(t, tt.isBannedFunc(tt.value)) + + // Ban the item + err := tt.banFunc(tt.value) + assert.NoError(t, err) + + // Should be banned + assert.True(t, tt.isBannedFunc(tt.value)) + + // Should appear in list + items, err := tt.listFunc() + assert.NoError(t, err) + assert.Contains(t, items, tt.value) + + // Unban the item + err = tt.unbanFunc(tt.value) + assert.NoError(t, err) + + // Should not be banned + assert.False(t, tt.isBannedFunc(tt.value)) + }) + } +}
\ No newline at end of file diff --git a/internal/mobius/transaction_handlers.go b/internal/mobius/transaction_handlers.go index eacf4c1..ebf18e6 100644 --- a/internal/mobius/transaction_handlers.go +++ b/internal/mobius/transaction_handlers.go @@ -8,7 +8,6 @@ import ( "fmt" "io" "math/big" - "net" "os" "path" "strings" @@ -1047,25 +1046,28 @@ func HandleTranAgreed(cc *hotline.ClientConn, t *hotline.Transaction) (res []hot } } + login := cc.Account.Login + ip := cc.IP() + if cc.Server.Redis != nil { - login := cc.Account.Login - ip, _, _ := net.SplitHostPort(cc.RemoteAddr) // Remove old entry (login::ip) - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+"::"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+"::"+ip) // Add new entry with login, nickname, ip - cc.Server.Redis.SAdd(context.Background(), "mobius:online", login+":"+string(cc.UserName)+":"+ip) - // Ban check for nickname - bannedNick, _ := cc.Server.Redis.SIsMember(context.Background(), "mobius:banned:nicknames", string(cc.UserName)).Result() - if bannedNick { + cc.Server.Redis.SAdd(context.Background(), hotline.RedisKeyOnline, login+":"+string(cc.UserName)+":"+ip) + } + + // Ban check for nickname + if cc.Server.BanList != nil && cc.Server.BanList.IsNicknameBanned(string(cc.UserName)) { + if cc.Server.Redis != nil { // Remove all possible online entries for this login and IP - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+"::"+ip) - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+":"+string(cc.UserName)+":"+ip) - // If we track the previous nickname, remove that too: - // cc.Server.Redis.SRem(context.Background(), "mobius:online", login+":"+oldNickname+":"+ip) - cc.Server.Redis.SAdd(context.Background(), "mobius:banned:ips", ip) - cc.Disconnect() - return res + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+"::"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+":"+string(cc.UserName)+":"+ip) } + if err := cc.Server.BanList.Add(ip, nil); err != nil { + cc.Logger.Error("Failed to ban IP for banned nickname", "ip", ip, "err", err) + } + cc.Disconnect() + return res } cc.Icon = t.GetField(hotline.FieldUserIconID).Data @@ -1191,7 +1193,7 @@ func HandleDisconnectUser(cc *hotline.ClientConn, t *hotline.Transaction) (res [ )) banUntil := time.Now().Add(hotline.BanDuration) - ip, _, _ := net.SplitHostPort(clientConn.RemoteAddr) + ip := clientConn.IP() err := cc.Server.BanList.Add(ip, &banUntil) if err != nil { @@ -1209,7 +1211,7 @@ func HandleDisconnectUser(cc *hotline.ClientConn, t *hotline.Transaction) (res [ hotline.NewField(hotline.FieldChatOptions, []byte{0, 0}), )) - ip, _, _ := net.SplitHostPort(clientConn.RemoteAddr) + ip := clientConn.IP() err := cc.Server.BanList.Add(ip, nil) if err != nil { @@ -1802,29 +1804,35 @@ func HandleSetClientUserInfo(cc *hotline.ClientConn, t *hotline.Transaction) (re oldNickname := string(cc.UserName) newNickname := string(t.GetField(hotline.FieldUserName).Data) cc.UserName = t.GetField(hotline.FieldUserName).Data + + login := cc.Account.Login + ip := cc.IP() + if cc.Server.Redis != nil { - login := cc.Account.Login - ip, _, _ := net.SplitHostPort(cc.RemoteAddr) // Remove old entry (login:oldnickname:ip) and (login::ip) - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+"::"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+"::"+ip) if oldNickname != "" { - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+":"+oldNickname+":"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+":"+oldNickname+":"+ip) } // Add new entry - cc.Server.Redis.SAdd(context.Background(), "mobius:online", login+":"+newNickname+":"+ip) - // Ban check for nickname - bannedNick, _ := cc.Server.Redis.SIsMember(context.Background(), "mobius:banned:nicknames", newNickname).Result() - if bannedNick { + cc.Server.Redis.SAdd(context.Background(), hotline.RedisKeyOnline, login+":"+newNickname+":"+ip) + } + + // Ban check for nickname + if cc.Server.BanList != nil && cc.Server.BanList.IsNicknameBanned(newNickname) { + if cc.Server.Redis != nil { // Remove all possible online entries for this login and IP - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+"::"+ip) - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+":"+newNickname+":"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+"::"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+":"+newNickname+":"+ip) if oldNickname != "" { - cc.Server.Redis.SRem(context.Background(), "mobius:online", login+":"+oldNickname+":"+ip) + cc.Server.Redis.SRem(context.Background(), hotline.RedisKeyOnline, login+":"+oldNickname+":"+ip) } - cc.Server.Redis.SAdd(context.Background(), "mobius:banned:ips", ip) - cc.Disconnect() - return res } + if err := cc.Server.BanList.Add(ip, nil); err != nil { + cc.Logger.Error("Failed to ban IP for banned nickname", "ip", ip, "err", err) + } + cc.Disconnect() + return res } } |