Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -6,4 +6,5 @@
/Downloads/
/Rshell_darwin_amd64
/Rshell_darwin_arm64
/.gitignore
/.gitignore
./Rshell.exe
Binary file added assets/image-20260420125544426.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
124 changes: 94 additions & 30 deletions pkg/api/forward.go
Original file line number Diff line number Diff line change
@@ -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"

Expand All @@ -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)
Expand All @@ -58,7 +61,6 @@ func ForwardConnect(c *gin.Context) {
return
}

// 返回连接信息
c.JSON(http.StatusOK, gin.H{
"status": 200,
"data": gin.H{
Expand All @@ -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)
Expand All @@ -108,7 +105,6 @@ func handleTCPForward(address, proxy string, c *gin.Context) {
return
}

// 返回连接信息
c.JSON(http.StatusOK, gin.H{
"status": 200,
"data": gin.H{
Expand All @@ -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)
Expand Down
62 changes: 46 additions & 16 deletions pkg/api/users.go
Original file line number Diff line number Diff line change
@@ -1,5 +1,11 @@
package api

/*
修改说明:
1. ChangePasswordHandler 添加用户存在判断。
2. LoginHandler 添加用户存在判断。
*/

import (
"Rshell/pkg/common"
"Rshell/pkg/database"
Expand All @@ -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,
}})
}

// 注销处理函数
Expand All @@ -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"})
Expand Down
43 changes: 38 additions & 5 deletions pkg/database/sqlite3.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,8 @@ package database

import (
"Rshell/pkg/logger"
"crypto/rand"
"fmt"
"log"
"os"
"path/filepath"
"reflect"
Expand Down Expand Up @@ -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)
}

// 获取程序所在目录
Expand All @@ -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,
Expand All @@ -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)
Expand Down
Loading
Loading