diff options
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 } } |