diff --git a/.gitignore b/.gitignore index 7937caf..c5b9435 100644 --- a/.gitignore +++ b/.gitignore @@ -6,4 +6,5 @@ /Downloads/ /Rshell_darwin_amd64 /Rshell_darwin_arm64 -/.gitignore \ No newline at end of file +/.gitignore +./Rshell.exe \ No newline at end of file diff --git a/assets/image-20260420125544426.png b/assets/image-20260420125544426.png new file mode 100644 index 0000000..1865418 Binary files /dev/null and b/assets/image-20260420125544426.png differ diff --git a/pkg/api/forward.go b/pkg/api/forward.go index 11af346..ba46cb7 100644 --- a/pkg/api/forward.go +++ b/pkg/api/forward.go @@ -1,12 +1,21 @@ -// api/forward.go package api +/* +修改说明: +1. api/forward-connection 添加目标地址规范化与校验逻辑。 +2. 拒绝 localhost、回环地址、私网地址、未指定地址、组播地址、链路本地地址。 +3. 对域名先做解析,只要解析结果指向受限地址就拒绝连接。 + +*/ + import ( "Rshell/pkg/connection/tcp" "Rshell/pkg/connection/websocket" "Rshell/pkg/logger" "fmt" + "net" "net/http" + "strconv" "strings" "time" @@ -26,28 +35,22 @@ func ForwardConnect(c *gin.Context) { switch forward.Type { case "websocket": - // 验证地址格式 - if forward.Address == "" { - c.JSON(http.StatusOK, gin.H{"status": 400, "data": "address is required"}) + safeAddress, err := normalizeForwardAddress(forward.Address, "") + if err != nil { + c.JSON(http.StatusOK, gin.H{"status": 400, "data": err.Error()}) return } - // 确保地址有正确的协议前缀 - - // 配置正向连接 config := &websocket.ForwardConfig{ - ServerURL: "ws://" + forward.Address + "/ws", - Socks5Proxy: forward.Proxy, // 使用传入的代理地址 + ServerURL: "ws://" + safeAddress + "/ws", + Socks5Proxy: forward.Proxy, Timeout: 30 * time.Second, MaxRetries: 5, RetryDelay: 10 * time.Second, Reconnect: true, - Headers: map[string]string{ - //"User-Agent": "Rshell-Forward-Client/1.0", - }, + Headers: map[string]string{}, } - // 启动正向客户端 client, err := websocket.StartForwardClient(config) if err != nil { logger.Error("Failed to start forward client:", err) @@ -58,7 +61,6 @@ func ForwardConnect(c *gin.Context) { return } - // 返回连接信息 c.JSON(http.StatusOK, gin.H{ "status": 200, "data": gin.H{ @@ -69,35 +71,30 @@ func ForwardConnect(c *gin.Context) { }, }) - // 异步监控连接状态 go monitorForwardConnection(client.UID, "websocket") case "tcp": - // TCP正向连接 handleTCPForward(forward.Address, forward.Proxy, c) default: c.JSON(http.StatusOK, gin.H{"status": 400, "data": "unsupported connection type"}) } } -// handleTCPForward 处理TCP正向连接 func handleTCPForward(address, proxy string, c *gin.Context) { - // 验证TCP地址格式(host:port) - if !strings.Contains(address, ":") { - // 如果没有端口,添加默认端口 - address = address + ":8080" + safeAddress, err := normalizeForwardAddress(address, "8080") + if err != nil { + c.JSON(http.StatusOK, gin.H{"status": 400, "data": err.Error()}) + return } - // 配置TCP正向连接 config := &tcp.TCPForwardConfig{ - ServerAddress: address, - Socks5Proxy: proxy, // 使用传入的代理地址 + ServerAddress: safeAddress, + Socks5Proxy: proxy, Timeout: 30 * time.Second, MaxRetries: 5, RetryDelay: 10 * time.Second, Reconnect: true, } - // 启动TCP正向客户端 client, err := tcp.StartTCPForwardClient(config) if err != nil { logger.Error("Failed to start TCP forward client:", err) @@ -108,7 +105,6 @@ func handleTCPForward(address, proxy string, c *gin.Context) { return } - // 返回连接信息 c.JSON(http.StatusOK, gin.H{ "status": 200, "data": gin.H{ @@ -120,16 +116,84 @@ func handleTCPForward(address, proxy string, c *gin.Context) { }, }) - // 异步监控连接状态 go monitorForwardConnection(client.UID, "tcp") } -// monitorForwardConnection 监控正向连接状态 +func normalizeForwardAddress(address, defaultPort string) (string, error) { + address = strings.TrimSpace(address) + if address == "" { + return "", fmt.Errorf("address is required") + } + + if strings.ContainsAny(address, "/?#") { + return "", fmt.Errorf("invalid address format") + } + + host, port, err := net.SplitHostPort(address) + if err != nil { + if defaultPort == "" || !strings.Contains(err.Error(), "missing port in address") { + return "", fmt.Errorf("invalid address format, expected host:port") + } + host = address + port = defaultPort + } + + host = strings.Trim(host, "[]") + if host == "" { + return "", fmt.Errorf("host is required") + } + + portNum, err := strconv.Atoi(port) + if err != nil || portNum < 1 || portNum > 65535 { + return "", fmt.Errorf("invalid port") + } + + if err := validateForwardHost(host); err != nil { + return "", err + } + + return net.JoinHostPort(host, strconv.Itoa(portNum)), nil +} + +func validateForwardHost(host string) error { + normalizedHost := strings.ToLower(strings.TrimSpace(host)) + if normalizedHost == "localhost" { + return fmt.Errorf("restricted target host") + } + + if ip := net.ParseIP(normalizedHost); ip != nil { + if isRestrictedForwardIP(ip) { + return fmt.Errorf("restricted target address") + } + return nil + } + + ips, err := net.LookupIP(normalizedHost) + if err != nil || len(ips) == 0 { + return fmt.Errorf("failed to resolve target host") + } + + for _, ip := range ips { + if isRestrictedForwardIP(ip) { + return fmt.Errorf("target resolves to restricted address") + } + } + + return nil +} + +func isRestrictedForwardIP(ip net.IP) bool { + return ip.IsLoopback() || + ip.IsPrivate() || + ip.IsUnspecified() || + ip.IsMulticast() || + ip.IsLinkLocalUnicast() || + ip.IsLinkLocalMulticast() +} + func monitorForwardConnection(uid, connType string) { - // 等待一段时间检查连接状态 time.Sleep(3 * time.Second) - // 根据连接类型检查是否还是临时UID if connType == "websocket" { if strings.HasPrefix(uid, "temp_") { logger.Warn("WebSocket forward connection still using temporary UID after 3 seconds:", uid) diff --git a/pkg/api/users.go b/pkg/api/users.go index 9c738c9..d111d71 100644 --- a/pkg/api/users.go +++ b/pkg/api/users.go @@ -1,5 +1,11 @@ package api +/* +修改说明: +1. ChangePasswordHandler 添加用户存在判断。 +2. LoginHandler 添加用户存在判断。 +*/ + import ( "Rshell/pkg/common" "Rshell/pkg/database" @@ -19,23 +25,29 @@ func LoginHandler(c *gin.Context) { return } - // 假设用户名和密码验证成功 var users database.Users - if database.Engine.Where("username = ?", loginData.Username).Get(&users); users.Password == loginData.Password { - token, err := common.GenerateJWT(loginData.Username) - if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": "Could not generate token"}) - return - } - c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{ - "token": token, - "permissions": 1, // 示例:1表示管理员权限 - "refresh": "mock-refresh-token", - "username": loginData.Username, - }}) - } else { + has, err := database.Engine.Where("username = ?", loginData.Username).Get(&users) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Database error"}) + return + } + + if !has || users.Password != loginData.Password { c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"}) + return } + + token, err := common.GenerateJWT(loginData.Username) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Could not generate token"}) + return + } + c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{ + "token": token, + "permissions": 1, // 示例:1表示管理员权限 + "refresh": "mock-refresh-token", + "username": loginData.Username, + }}) } // 注销处理函数 @@ -52,15 +64,33 @@ func ChangePasswordHandler(c *gin.Context) { } if err := c.ShouldBind(&passwordData); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": "Invalid input"}) + return } // 处理密码修改逻辑 if passwordData.OldPassword != passwordData.NewPassword { username := c.MustGet("username").(string) var users database.Users - if database.Engine.Where("username = ?", username).Get(&users); users.Password == passwordData.OldPassword { + has, err := database.Engine.Where("username = ?", username).Get(&users) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Database error"}) + return + } + if !has { + c.JSON(http.StatusOK, gin.H{"code": 400, "message": "Password changed failed"}) + return + } + if users.Password == passwordData.OldPassword { users.Password = passwordData.NewPassword - database.Engine.Where("username = ?", username).Update(&users) + affected, err := database.Engine.Where("username = ?", username).Cols("password").Update(&users) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Database error"}) + return + } + if affected != 1 { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Password update did not persist"}) + return + } c.JSON(http.StatusOK, gin.H{"code": 200, "message": "Password changed successfully"}) } else { c.JSON(http.StatusOK, gin.H{"code": 400, "message": "Password changed failed"}) diff --git a/pkg/database/sqlite3.go b/pkg/database/sqlite3.go index 396201b..371030f 100644 --- a/pkg/database/sqlite3.go +++ b/pkg/database/sqlite3.go @@ -2,8 +2,8 @@ package database import ( "Rshell/pkg/logger" + "crypto/rand" "fmt" - "log" "os" "path/filepath" "reflect" @@ -88,12 +88,32 @@ type Key struct { PrivateKey string } +func generateInitialAdminPassword(length int) (string, error) { + if length <= 0 { + length = 20 + } + + const alphabet = "ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnopqrstuvwxyz23456789!@#$%^&*_-+=" + buf := make([]byte, length) + randomBytes := make([]byte, length) + + if _, err := rand.Read(randomBytes); err != nil { + return "", fmt.Errorf("failed to generate random password: %v", err) + } + + for i := range buf { + buf[i] = alphabet[int(randomBytes[i])%len(alphabet)] + } + + return string(buf), nil +} + func ConnectDateBase() { var err error // 获取当前程序所在目录 exePath, err := os.Executable() if err != nil { - log.Fatalf("获取程序路径失败: %v", err) + logger.Fatalf("获取程序路径失败: %v", err) } // 获取程序所在目录 @@ -104,19 +124,28 @@ func ConnectDateBase() { Engine, err = xorm.NewEngine("sqlite", DatabaseName) if err != nil { - log.Fatalf("连接sqlite数据库失败: %v", err) + logger.Fatalf("连接sqlite数据库失败: %v", err) } err = Engine.Sync2(new(Users), new(Clients), new(Notes), new(Shell), new(Downloads), new(Listener), new(WebDelivery), new(Socks5), new(Settings), new(Key)) if err != nil { - log.Fatalf("初始化数据库失败: %v", err) + logger.Fatalf("初始化数据库失败: %v", err) } var user Users exists, err := Engine.Where("username = ?", "admin").Get(&user) + if err != nil { + logger.Fatalf("查询 admin 用户失败: %v", err) + } if !exists { + initialPassword, err := generateInitialAdminPassword(20) + if err != nil { + logger.Fatalf("生成初始 admin 密码失败: %v", err) + } + // 如果不存在 admin 用户,插入默认的 admin 用户 + // 当 admin 用户不存在时,改为生成随机初始密码并写入数据库。 defaultUser := &Users{ Username: "admin", - Password: "admin123", + Password: initialPassword, //Email: "admin@example.com", //Phone: "1234567890", //Permissions: 1, @@ -127,6 +156,10 @@ func ConnectDateBase() { logger.Error(fmt.Sprintf("插入默认 admin 用户失败: %v", err)) os.Exit(0) } + + logger.Warn("admin user not found,init......") + logger.Warnf("account: %s", "admin") + logger.Warnf("password: %s", initialPassword) } var Setting Settings exists, err = Engine.Where("name=?", "wecom").Get(&Setting) diff --git a/pkg/middlewares/auth.go b/pkg/middlewares/auth.go index f1eb838..7c8ff14 100644 --- a/pkg/middlewares/auth.go +++ b/pkg/middlewares/auth.go @@ -1,5 +1,11 @@ package middlewares +/* +修改说明: +1. BasicAuthMiddleware 添加用户存在判断。 +2. Authorization2 添加合法性判断。 +*/ + import ( "Rshell/pkg/common" "Rshell/pkg/database" @@ -11,12 +17,10 @@ import ( ) func BasicAuthMiddleware() gin.HandlerFunc { - return func(c *gin.Context) { authHeader := c.Request.Header.Get("Authorization") if authHeader == "" || !strings.HasPrefix(authHeader, "Basic ") { - // 返回WWW-Authenticate头,触发浏览器的弹框 c.Header("WWW-Authenticate", `Basic realm="Restricted"`) c.AbortWithStatus(http.StatusUnauthorized) return @@ -38,9 +42,15 @@ func BasicAuthMiddleware() gin.HandlerFunc { } user, pass := credParts[0], credParts[1] - var user_pass database.Users - database.Engine.Where("username = ?", user).Get(&user_pass) - if user_pass.Password != pass || user_pass.Password == "" { + var userPass database.Users + has, err := database.Engine.Where("username = ?", user).Get(&userPass) + if err != nil || !has { + c.Header("WWW-Authenticate", `Basic realm="Restricted"`) + c.AbortWithStatus(http.StatusUnauthorized) + return + } + + if userPass.Password != pass || userPass.Password == "" { c.Header("WWW-Authenticate", `Basic realm="Restricted"`) c.AbortWithStatus(http.StatusUnauthorized) return @@ -51,20 +61,29 @@ func BasicAuthMiddleware() gin.HandlerFunc { } } -// JWT 验证中间件 +// AuthMiddleware validates JWT from Authorization2. func AuthMiddleware() gin.HandlerFunc { return func(c *gin.Context) { - if c.GetHeader("Authorization2") == "" { + authHeader := strings.TrimSpace(c.GetHeader("Authorization2")) + if authHeader == "" { c.JSON(http.StatusUnauthorized, gin.H{"error": "Token required"}) c.Abort() return } - tokenString := c.GetHeader("Authorization2")[len("Bearer "):] + + if len(authHeader) < len("Bearer ") || !strings.EqualFold(authHeader[:len("Bearer ")], "Bearer ") { + c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid token format"}) + c.Abort() + return + } + + tokenString := strings.TrimSpace(authHeader[len("Bearer "):]) if tokenString == "" { c.JSON(http.StatusUnauthorized, gin.H{"error": "Token required"}) c.Abort() return } + claims, err := common.ValidateJWT(tokenString) if err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid token"}) diff --git a/readme.md b/readme.md index 108b272..47fc4c2 100644 --- a/readme.md +++ b/readme.md @@ -8,11 +8,18 @@ Rshell是一款开源的golang编写的支持多平台的C2框架,旨在帮助 通过-p参数指定端口(默认端口8089)并运行: -``` +```bash ./Rshell -p 8089 ``` -**默认账号密码:admin/admin123** +## 账号密码 + +**账号:admin** + +**密码:首次运行后随机生成** + +![image-20260420125544426](./assets/image-20260420125544426.png) + ![image-20260205112117018](./assets/image-20260205112117018.png)