aboutsummaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/mobius/api.go118
-rw-r--r--internal/mobius/ban.go171
-rw-r--r--internal/mobius/ban_test.go152
-rw-r--r--internal/mobius/redis_ban.go200
-rw-r--r--internal/mobius/redis_ban_test.go222
-rw-r--r--internal/mobius/transaction_handlers.go70
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
}
}