优化连接逻辑,更新readme

This commit is contained in:
Mob2003
2023-04-22 19:46:52 +08:00
parent 5a67527907
commit a8dcc3871d
21 changed files with 1310 additions and 1443 deletions
+781 -4
View File
@@ -1,11 +1,27 @@
package server
import (
"bytes"
"cert"
"encoding/json"
"fmt"
"io"
"io/ioutil"
"os"
"rakshasa/aes"
"rakshasa/common"
"regexp"
"runtime"
"strconv"
"strings"
"sync/atomic"
"time"
"github.com/abiosoft/readline"
"github.com/google/uuid"
"github.com/luyu6056/ishell"
"golang.org/x/text/encoding/simplifiedchinese"
"golang.org/x/text/transform"
)
var rootCli = cliInit()
@@ -99,9 +115,9 @@ func cliInit() *ishell.Shell {
c.Println("参数不对")
return
}
n, ok := nodeMap[c.Args[0]]
n, ok := nodeMap.Load(c.Args[0])
if ok {
n.Delete("")
n.(*node).Delete("")
}
},
@@ -116,9 +132,9 @@ func cliInit() *ishell.Shell {
c.Println("参数不对")
return
}
n, ok := nodeMap[c.Args[0]]
n, ok := nodeMap.Load(c.Args[0])
if ok {
n.Close("debug关闭")
n.(*node).Close("debug关闭")
}
},
@@ -126,3 +142,764 @@ func cliInit() *ishell.Shell {
}
return shell
}
func init() {
configShell := cliInit()
configShell.SetPrompt("rakshasa\\config>")
configShell.AddCmd(&ishell.Cmd{
Name: "info",
Help: "打印当前配置",
Func: func(c *ishell.Context) {
c.Println("当前节点", currentNode.uuid)
c.Println("上级节点地址", currentConfig.DstNode)
c.Println("通讯密码", currentConfig.Password)
c.Println("监听端口", currentConfig.Port)
c.Println("监听IP", currentConfig.ListenIp)
c.Println("禁止额外连接", currentConfig.Limit)
c.Println("配置文件名", currentConfig.FileName)
if currentConfig.FileSave {
c.Println("当前配置:已写入文件")
} else {
c.Println("当前配置:未写入文件")
}
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "save",
Help: "保存文件",
Func: func(c *ishell.Context) {
if err := ConfigSave(); err == nil {
c.Println("写入成功")
} else {
c.Println("保存失败", err.Error())
}
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "d",
Help: "修改上级节点地址,格式为 ip:端口 多个节点以,隔开 注意:不会立刻连接设置节点, 当发生 节点掉线重连 时候会连接该地址",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,格式为 ip:端口 多个节点以,隔开 如 d 192.168.1.1:8883,192.168.1.2:8883")
return
}
dstNode, err := common.ResolveTCPAddr(c.Args[0])
if err != nil {
c.Println("参数错误,格式为 ip:端口 多个节点以,隔开 如 d 192.168.1.1:8883,192.168.1.2:8883")
return
}
currentConfig.DstNode = dstNode
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "password",
Help: "修改通讯密码,立即生效",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,格式为 password \"123456\"")
return
}
c.Println(c.Args)
currentConfig.Password = c.Args[0]
currentConfig.FileSave = false
aes.Key = aes.MD5_B(currentConfig.Password + string(cert.RsaPrivateKey[:16]))
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "port",
Help: "修改监听端口,立即生效",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,格式为 port 8883")
return
}
port, _ := strconv.Atoi(c.Args[0])
if port <= 0 || port > 65535 {
c.Println("参数错误,端口范围是1-65535")
return
}
c.Println("正在关闭server监听")
if currentNode.listen != nil {
currentNode.listen.Close()
currentNode.listen = nil
}
currentConfig.Port = port
currentNode.port = port
currentConfig.FileSave = false
if err := StartServer(fmt.Sprintf(":%d", currentConfig.Port)); err != nil {
c.Printf("启动节点失败 %v, 请重新修改监听端口", currentConfig.Port)
}
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "ip",
Help: "修改本节点连接ip,当其他节点进行额外连接时候,优先使用此ip连接",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
currentConfig.ListenIp = c.Args[0]
currentNode.mainIp = currentConfig.ListenIp
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "limit",
Help: "修改本节点Limit设置,使用方法 limit true",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
currentConfig.Limit = c.Args[0] == "true"
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "f",
Help: "修改配置文件名,使用方法 f config.yaml",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
currentConfig.FileName = c.Args[0]
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "uuid",
Help: "修改本节点UUID设置,使用方法uuid 字串符",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
if id, err := uuid.Parse(c.Args[0]); err == nil {
nodeMap.Delete(currentConfig.UUID)
currentConfig.UUID = id.String()
nodeMap.Store(currentConfig.UUID, currentNode)
currentConfig.FileSave = false
SetConfig(currentConfig)
} else {
c.Println("输入的uuid不是合法的uuid,建议使用xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx")
}
},
})
rootCli.AddCmd(&ishell.Cmd{
Name: "config",
Help: "配置管理",
Func: func(c *ishell.Context) {
configShell.Run()
},
})
remoteShell := cliInit()
remoteShell.SetPrompt("rakshasa\\remoteshell>")
fileShell := cliInit()
remoteShell.AddCmd(&ishell.Cmd{
Name: "file",
Help: "连到节点进行文件管理,参数为id或者uuid",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
workN, _ := getNode(c.Args[0])
if workN == nil {
c.Println("无法连接节点", c.Args[0])
return
}
if workN != nil {
fileShell.Set("node", workN)
result := make(chan interface{}, 1)
id := workN.storeQuery(result)
workN.Write(common.CMD_PWD, id, []byte(cert.RSAEncrypterByPriv(currentConfig.Password)))
select {
case pwd := <-result:
workN.deleteQuery(id)
pwd = strings.ReplaceAll(pwd.(string), "\\", "/")
fileShell.Set("pwd", pwd)
fileShell.SetPrompt(workN.uuid + " " + pwd.(string) + ">")
fileShell.Run()
case <-time.After(common.CMD_TIMEOUT):
workN.deleteQuery(id)
c.Println("连接", c.Args[0], "超时")
}
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "dir",
Help: "打印当前目录文件",
Func: func(c *ishell.Context) {
pwd := fileShell.Get("pwd")
n := c.Get("node").(*node)
resChan := make(chan interface{}, 1)
id := n.storeQuery(resChan)
n.Write(common.CMD_DIR, id, []byte(cert.RSAEncrypterByPriv(pwd.(string))))
select {
case res := <-resChan:
n.deleteQuery(id)
c.Println(res)
case <-time.After(common.CMD_TIMEOUT):
n.deleteQuery(id)
c.Println("dir time out")
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "cd",
Help: "切换工作目录",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
dir := c.Args[0]
pwd := fileShell.Get("pwd").(string)
n := c.Get("node").(*node)
if strings.Contains(dir, ":/") || dir[0] == '/' || dir == "~" {
pwd = dir
} else {
pwd += "/" + dir
pwd = strings.TrimRight(getRealPath(pwd), "/")
}
resChan := make(chan interface{}, 1)
id := n.storeQuery(resChan)
n.Write(common.CMD_CD, id, []byte(cert.RSAEncrypterByPriv(pwd)))
select {
case res := <-resChan:
n.deleteQuery(id)
if err, ok := res.(error); ok {
c.Println(err.Error())
} else {
pwd = res.(string)
fileShell.Set("pwd", pwd)
c.SetPrompt(n.uuid + " " + pwd + ">")
}
case <-time.After(common.CMD_TIMEOUT):
n.deleteQuery(id)
c.Println("dir time out")
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "upload",
Help: "上传文件 ,upload 本地文件 远程目录(为空传到工作目录)",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 && len(c.Args) != 2 {
c.Println("参数错误")
return
}
s, err := os.Stat(c.Args[0])
if err != nil {
c.Println("打开本地文件", c.Args[0], "错误 ", err)
return
}
f, err := os.Open(c.Args[0])
if err != nil {
c.Println("打开本地文件", c.Args[0], "错误 ", err)
return
}
defer f.Close()
pwd := fileShell.Get("pwd").(string) + "/"
n := c.Get("node").(*node)
if len(c.Args) == 2 {
pwd = c.Args[1]
}
pwd = strings.ReplaceAll(pwd, "\\", "/")
c.Args[0] = strings.ReplaceAll(c.Args[0], "\\", "/")
i := strings.LastIndex(c.Args[0], "/")
if i == -1 {
i = 0
}
if pwd[len(pwd)-1] == '/' {
pwd += c.Args[0][i:]
}
i = strings.LastIndex(pwd, "/")
if i == -1 {
i = 0
}
filename := pwd[i+1:]
dir := pwd[:i]
dir = strings.TrimRight(getRealPath(dir), "/") + "/"
pwd = dir + filename
resChan := make(chan interface{}, 9999) //避免收消息阻塞
filereadChan := make(chan []byte, 10)
upload := func() {
for i := 0; i < 10; i++ {
buf := make([]byte, common.MAX_PACKAGE-len(pwd)-9)
n, err := f.Read(buf)
if err != nil {
if err == io.EOF {
return
}
resChan <- err
c.Println("读取文件", c.Args[0], "错误", err)
return
}
filereadChan <- buf[:n]
}
}
offset := 0
be := len(pwd) + 1
id := n.storeQuery(resChan)
defer n.deleteQuery(id)
b := []byte(pwd)
b = append(b, 0, 0, 0, 0, 0, 0, 0, 0, 0)
c.ProgressBar().Start()
go upload()
var resnum int
for {
select {
case data := <-filereadChan:
b[be] = byte(offset)
b[be+1] = byte(offset >> 8)
b[be+2] = byte(offset >> 16)
b[be+3] = byte(offset >> 24)
b[be+4] = byte(offset >> 32)
b[be+5] = byte(offset >> 40)
b[be+6] = byte(offset >> 48)
b[be+7] = byte(offset >> 56)
offset += len(data)
n.Write(common.CMD_UPLOAD, id, cert.RSAEncrypterByPrivByte(append(b, data...)))
case res := <-resChan:
switch v := res.(type) {
case error:
c.ProgressBar().Stop()
c.Println("上传失败", res)
return
case int64:
resnum++
i := v * 100 / s.Size()
c.ProgressBar().Suffix(fmt.Sprint(" ", i, "%"))
c.ProgressBar().Progress(int(i))
if v == s.Size() {
c.ProgressBar().Stop()
c.Println(c.Args[0], "上传成功")
return
}
if resnum >= 5 {
go upload()
resnum -= 10
}
default:
c.Println("协议错误")
return
}
case <-time.After(common.CMD_TIMEOUT):
c.ProgressBar().Stop()
c.Println("upload time out")
return
}
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "download",
Help: "下载文件 download 远程文件 本地目录(为空本地执行目录)",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 && len(c.Args) != 2 {
c.Println("参数错误")
return
}
pwd := fileShell.Get("pwd").(string)
n := c.Get("node").(*node)
file := c.Args[0]
file = strings.ReplaceAll(file, "\\", "/")
if strings.Contains(file, ":/") || file[0] == '/' {
pwd = file
} else {
pwd += "/" + file
}
i := strings.LastIndex(pwd, "/")
if i == -1 {
i = 0
}
filename := pwd[i+1:]
dir := pwd[:i]
dir = strings.TrimRight(getRealPath(dir), "/") + "/"
mydir, err := os.Getwd()
local := "./" + filename
if err == nil {
local = mydir + "/" + filename
}
if len(c.Args) == 2 {
s, err := os.Stat(c.Args[1])
if err == nil {
if s.IsDir() {
local = strings.TrimRight(c.Args[1], "/") + "/" + filename
} else {
local = c.Args[1]
}
} else {
local = c.Args[1]
}
}
pwd = dir + filename
result := make(chan interface{}, 999)
id := n.storeQuery(result)
defer n.deleteQuery(id)
b := []byte(pwd)
b = append(b, []byte{0, 0, 0, 0, 0, 0, 0, 0, 0}...)
total := int64(-1)
be := len(pwd) + 1
b[be] = byte(total)
b[be+1] = byte(total >> 8)
b[be+2] = byte(total >> 16)
b[be+3] = byte(total >> 24)
b[be+4] = byte(total >> 32)
b[be+5] = byte(total >> 40)
b[be+6] = byte(total >> 48)
b[be+7] = byte(total >> 56)
n.Write(common.CMD_DOWNLOAD, id, cert.RSAEncrypterByPrivByte(b))
c.ProgressBar().Start()
size := int64(0)
resnum := 0
total = 0
var f *os.File
for {
select {
case res := <-result:
switch v := res.(type) {
case error:
c.ProgressBar().Stop()
c.Println("下载失败", res)
return
case int64:
var err error
size = v
f, err = os.OpenFile(local, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0666)
if err != nil {
c.Println("本地文件 ", local, "写入失败", err.Error())
return
}
defer f.Close()
case []byte:
if f == nil {
c.Println("本地文件 ", local, "不可写入")
return
}
resnum++
num, err := f.Write(v)
if err != nil {
c.Println("本地文件 ", local, "写入失败", err.Error())
return
}
if num != len(v) {
c.Println("本地文件 ", local, "写入失败,写入量不符")
return
}
total += int64(num)
i := total * 100 / size
c.ProgressBar().Suffix(fmt.Sprint(" ", i, "%"))
c.ProgressBar().Progress(int(i))
if total == size {
c.ProgressBar().Stop()
c.Println(c.Args[0], "下载成功 文件保存到", local)
return
}
if resnum == 10 {
resnum -= 10
b[be] = byte(total)
b[be+1] = byte(total >> 8)
b[be+2] = byte(total >> 16)
b[be+3] = byte(total >> 24)
b[be+4] = byte(total >> 32)
b[be+5] = byte(total >> 40)
b[be+6] = byte(total >> 48)
b[be+7] = byte(total >> 56)
n.Write(common.CMD_DOWNLOAD, id, cert.RSAEncrypterByPrivByte(b))
}
default:
c.Println("协议错误")
return
}
case <-time.After(common.CMD_TIMEOUT):
c.ProgressBar().Stop()
c.Println("upload time out")
return
}
}
},
})
remoteShell.AddCmd(&ishell.Cmd{
Name: "new",
Help: "与一个或者多个节点连接,使用方法 new ip:端口 多个地址以,间隔 如1080 127.0.0.1:1081,127.0.0.1:1082",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,使用方法 connect ip:端口")
return
}
for _, addr := range strings.Split(c.Args[0], ",") {
_, err := getNode(addr)
if err != nil {
c.Println("连接", addr, "失败", err)
return
}
}
},
})
remoteShell.AddCmd(&ishell.Cmd{
Name: "shell",
Help: "反弹shell 使用方法 shell id/uuid 启动参数 ,启动参数可为空,win默认启动cmd,linux默认启动bash 如 shell 1 powershell 。 shell 1 zsh",
Func: func(c *ishell.Context) {
if len(c.Args) < 1 {
c.Println("参数错误,例子 shell 1 powershell")
return
}
param := ""
if len(c.Args) == 2 {
param = c.Args[1]
}
n, _ := getNode(c.Args[0])
if n == nil {
c.Println("无法连接节点", c.Args[0])
return
}
res := make(chan interface{}, 999)
id := n.storeQuery(res)
defer n.deleteQuery(id)
p := StartCmdParam{
Param: param,
Size: common.GetSize(),
}
b, _ := json.Marshal(p)
n.Write(common.CMD_SHELL, id, cert.RSAEncrypterByPrivByte(b))
s := &remoteCmd{
cmd: nil,
stdin: nil,
inChan: make(chan []byte, 999),
translate: func(in []byte) ([]byte, error) { return in, nil },
pong: time.Now().Unix(),
}
select {
case i := <-res:
switch v := i.(type) {
case error:
c.Println("启动shell失败,错误", v.Error())
case []byte:
data := v
s.id = uint32(data[0]) | uint32(data[1])<<8 | uint32(data[2])<<16 | uint32(data[3])<<24
switch data[4] {
case 0: //windows
if string(data[len(data)-6:]) == string([]byte{32, 57, 51, 54, 13, 10}) { //活动代码页: 936
//gbk转utf8
s.translate = func(in []byte) ([]byte, error) {
reader := transform.NewReader(bytes.NewReader(in), simplifiedchinese.GBK.NewDecoder())
d, e := ioutil.ReadAll(reader)
if e != nil {
return nil, e
}
return d, nil
}
}
case 1: //linux
if runtime.GOOS == "windows" {
if !common.EnableTermVt {
s.translate = func(in []byte) ([]byte, error) {
if in[0] == 27 {
r, _ := regexp.Compile(`\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])`)
res := r.ReplaceAllString(string(in), "")
return []byte(res), nil
}
return in, nil
}
}
}
}
atomic.CompareAndSwapInt32(&s.cmdStatus, 0, 1)
}
case <-time.After(common.CMD_TIMEOUT):
c.Println("启动shell失败,超时")
return
}
n.shellMap.Store(s.id, s)
r, _ := readline.NewEx(&readline.Config{FuncIsTerminal: func() bool { return false }, ForcePrint: true})
defer func() {
n.shellMap.Delete(s.id)
atomic.StoreInt32(&s.cmdStatus, -1)
c.Println("请按回车键退出")
r.Close()
}()
go func() {
for {
switch s.cmdStatus {
case 1:
input, err := r.ReadlineEx()
if err != nil {
if err != readline.ErrInterrupt {
res <- err
return
}
if s.cmdStatus == 1 {
n.Write(common.CMD_SHELL_DATA, s.id, []byte{03})
}
}
if s.cmdStatus == 1 {
n.Write(common.CMD_SHELL_DATA, s.id, []byte(input+"\n"))
}
case 0:
time.Sleep(time.Millisecond * 100)
case -1:
return
}
}
}()
tick := time.NewTicker(common.CMD_TIMEOUT / 2)
for {
select {
case b := <-s.inChan:
s.pong = time.Now().Unix()
if len(b) > 0 {
b, err := s.translate(b)
if err != nil {
c.Println("shell 运行失败", err)
return
}
fmt.Print(string(b))
}
case v := <-res:
if err, ok := v.(error); ok {
if err.Error() != "退出shell" {
c.Println("运行shell", param, "失败", err)
}
} else {
c.Println("无法处理消息", v)
}
return
case <-tick.C:
s.ping = time.Now().Unix()
if s.ping-s.pong > int64(common.CMD_TIMEOUT/time.Second) {
c.Println("shell time out")
return
}
n.Write(common.CMD_SHELL_DATA, s.id, nil)
}
}
},
})
rootCli.AddCmd(&ishell.Cmd{
Name: "remoteshell",
Help: "远程shell",
Func: func(c *ishell.Context) {
remoteShell.Run()
},
})
}
// 打印节点
func printNodes(c *ishell.Context) {
l := clientLock.RLock()
defer l.RUnlock()
var list []*node
nodeMap.Range(func(key, value interface{}) bool {
n := value.(*node)
list = append(list, n)
return true
})
orderNode(list)
c.Println("ID UUID HostName GOOS IP listenIP")
c.Println("-----------------------------------------------------------------------------------------------------------------------------")
for k, n := range list {
n.id = k + 1
hostname := bytes.Repeat([]byte(" "), 22)
copy(hostname, n.hostName)
ip := bytes.Repeat([]byte(" "), 23)
if n.uuid == currentNode.uuid {
copy(ip, "(localhost)"+":"+strconv.Itoa(n.port))
} else {
copy(ip, n.addr+":"+strconv.Itoa(n.port))
}
listenip := n.mainIp
goos := bytes.Repeat([]byte(" "), 11)
copy(goos, n.goos)
c.Printf("%2d %s %s %s %s %s\n", n.id, n.uuid, hostname, goos, ip, listenip)
}
}
func getRealPath(path string) string {
path_s := strings.Split(path, "/")
realpath := []string{}
if len(path_s) == 0 {
return "error"
}
for _, value := range path_s {
if value == ".." {
k := len(realpath)
kk := k - 1
realpath = append(realpath[:kk], realpath[k:]...)
} else {
realpath = append(realpath, value)
}
}
return strings.Join(realpath, "/")
}
func printConn() {
connMap.Range(func(key, value interface{}) bool {
fmt.Println(key)
return true
})
}
+90 -110
View File
@@ -12,6 +12,7 @@ import (
"net/url"
"rakshasa/aes"
"rakshasa/common"
"runtime/debug"
"strconv"
"strings"
"sync"
@@ -34,13 +35,13 @@ type Conn struct {
node *node
nodeaddr string
//key string
remoteAddr string
inChan chan func()
OutChan chan []byte
close chan string
isClient bool
nodeConn *tls.Conn
regResult chan RegMsg
remoteAddr string
inChan chan func()
OutChan chan []byte
close chan string
isClient bool
nodeConn *tls.Conn
regResult chan RegMsg
}
type serverListen struct {
@@ -495,7 +496,6 @@ func (c *Conn) handlerNodeRead() {
if common.Debug {
fmt.Println("fromto", msg.From, msg.To, common.CmdToName[msg.CmdOpteion], int(lengbuf[0])+int(lengbuf[1])<<8)
}
if msg.To == common.NoneUUID.String() && c.node == nil {
c.inChan <- func() {
newNode := &node{
@@ -504,83 +504,77 @@ func (c *Conn) handlerNodeRead() {
newNode.do(msg)
}
} else if msg.To == currentNode.uuid {
func() {
l := clientLock.RLock()
v, ok := nodeMap[msg.From]
l.RUnlock()
if ok && v.port != 0 {
c.inChan <- func() {
v.do(msg)
v, ok := nodeMap.Load(msg.From)
if ok && v.(*node).port != 0 {
c.inChan <- func() {
v.(*node).do(msg)
}
} else {
if !ok {
newNode := &node{
uuid: msg.From,
conn: c,
waitMsg: []*common.Msg{msg},
}
} else {
l := clientLock.Lock()
v, ok := nodeMap[msg.From]
if !ok {
newNode := &node{
uuid: msg.From,
conn: c,
waitMsg: []*common.Msg{msg},
}
result := make(chan interface{}, 1)
id := newNode.storeQuery(result)
if common.Debug {
fmt.Printf("nodeMap1 %s %p \r\n", msg.From, newNode)
}
nodeMap[msg.From] = newNode
l.Unlock()
newNode.Write(common.CMD_GET_CURRENT_NODE, id, []byte{1}) //获取丢失节点的信息
go func() {
defer newNode.deleteQuery(id)
select {
case res := <-result:
if res == nil {
for _, m := range newNode.waitMsg {
c.inChan <- func() {
newNode.do(m)
}
result := make(chan interface{}, 1)
id := newNode.storeQuery(result)
nodeMap.Store(msg.From, newNode)
newNode.Write(common.CMD_GET_CURRENT_NODE, id, []byte{1}) //获取丢失节点的信息
go func() {
defer func() {
if err := recover(); err != nil {
fmt.Println(err)
debug.PrintStack()
}
newNode.deleteQuery(id)
}()
select {
case res := <-result:
if res == nil {
for _, m := range newNode.waitMsg {
c.inChan <- func() {
newNode.do(m)
}
}
case <-time.After(common.CMD_TIMEOUT):
newNode.Close("超时")
}
}()
case <-time.After(common.CMD_TIMEOUT):
newNode.Close("超时")
}
}()
} else {
if msg.CmdOpteion == common.CMD_GET_CURRENT_NODE_RESULT {
n := v.(*node)
var res chan interface{}
if _v, ok := v.loadQuery(msg.CmdId); !ok {
if _v, ok := n.loadQuery(msg.CmdId); !ok {
return
} else {
res = _v
}
var nmsg nodeInfo
err = json.Unmarshal(msg.CmdData, &nmsg)
if err != nil {
res <- err
return
}
v.hostName = cert.RSADecrypterStr(nmsg.HostName)
v.uuid = cert.RSADecrypterStr(nmsg.UUID)
if v.port, err = strconv.Atoi(cert.RSADecrypterStr(nmsg.Port)); err != nil {
v.port = -1
}
v.mainIp = cert.RSADecrypterStr(nmsg.MainIp)
v.goos = cert.RSADecrypterStr(nmsg.Goos)
res <- nil
} else {
v.waitMsg = append(v.waitMsg, msg)
var nmsg nodeInfo
err = json.Unmarshal(msg.CmdData, &nmsg)
if err != nil {
res <- err
return
}
n.hostName = cert.RSADecrypterStr(nmsg.HostName)
n.uuid = cert.RSADecrypterStr(nmsg.UUID)
if n.port, err = strconv.Atoi(cert.RSADecrypterStr(nmsg.Port)); err != nil {
n.port = -1
}
n.mainIp = cert.RSADecrypterStr(nmsg.MainIp)
n.goos = cert.RSADecrypterStr(nmsg.Goos)
res <- nil
} else {
v.(*node).waitMsg = append(v.(*node).waitMsg, msg)
}
l.Unlock()
}
}
}()
} else {
@@ -622,7 +616,6 @@ func (c *Conn) handlerNodeRead() {
func (c *Conn) handle() {
c.OutChan = make(chan []byte, 64)
c.inChan = make(chan func())
c.close = make(chan string, 999)
go func() {
@@ -648,43 +641,34 @@ func (c *Conn) handle() {
c.node.ping(0)
c.node.nextPingTime = time.Now().Unix() + 5
}
func() { //返回false则退出handle
connMap.Delete(c.remoteAddr)
l := clientLock.Lock()
defer func() {
l.Unlock()
}()
if atomic.CompareAndSwapInt32(&c.closeTag, 0, 1) {
if common.Debug {
fmt.Println(c.nodeConn.RemoteAddr().String(), "关闭原因", reason)
}
if c.nodeConn != nil {
if common.Debug {
fmt.Println("執行close1")
}
c.nodeConn.Close()
}
if c.node != nil {
c.node.Close(reason)
//移除上游连接
for i := len(upLevelNode) - 1; i >= 0; i-- {
n := upLevelNode[i]
if n.uuid == c.node.uuid {
upLevelNode = append(upLevelNode[:i], upLevelNode[i+1:]...)
}
}
if common.Debug {
fmt.Println("upLevelNode",len(upLevelNode))
}
}
connMap.Delete(c.remoteAddr)
if atomic.CompareAndSwapInt32(&c.closeTag, 0, 1) {
if common.Debug {
fmt.Println(c.nodeConn.RemoteAddr().String(), "关闭原因", reason)
}
return
}()
if c.nodeConn != nil {
if common.Debug {
fmt.Println("執行close1")
}
c.nodeConn.Close()
}
if c.node != nil {
c.node.Close(reason)
//移除上游连接
for i := len(upLevelNode) - 1; i >= 0; i-- {
n := upLevelNode[i]
if n.uuid == c.node.uuid {
upLevelNode = append(upLevelNode[:i], upLevelNode[i+1:]...)
}
}
if common.Debug {
fmt.Println("upLevelNode", len(upLevelNode))
}
}
}
return
}
}
@@ -715,12 +699,8 @@ func (c *Conn) reg() error {
return nil
}
func (c *Conn) WriteToUUID(msg *common.Msg) {
l := clientLock.RLock()
defer l.RUnlock()
if n, ok := nodeMap[msg.To]; ok {
n.WriteMsg(msg)
if n, ok := nodeMap.Load(msg.To); ok {
n.(*node).WriteMsg(msg)
}
}
@@ -733,10 +713,10 @@ func (c *Conn) tlsWrite(b []byte) error {
c.nodeConn.SetWriteDeadline(time.Now().Add(common.WRITE_DEADLINE))
n, err := c.nodeConn.Write(b)
if common.Debug {
if c.node!=nil{
if c.node != nil {
fmt.Println("writeto", c.node.uuid, n)
}else{
fmt.Println("writeto",common.NoneUUID, n)
} else {
fmt.Println("writeto", common.NoneUUID, n)
}
}
+151 -129
View File
@@ -12,20 +12,42 @@ import (
"math/rand"
"net"
"os"
"os/exec"
"rakshasa/common"
"runtime"
"runtime/debug"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"unsafe"
"github.com/creack/pty"
)
var (
shellMapLock sync.Mutex
)
type StartCmdParam struct {
Param string
Size *pty.Winsize
}
type remoteCmd struct {
cmdStatus int32
cmd *exec.Cmd
id uint32
stdin io.WriteCloser
inChan chan []byte
translate func(in []byte) ([]byte, error)
ping, pong int64
}
var (
currentNode = &node{}
clientLock = &lock{}
nodeMap = make(map[string]*node)
nodeMap = sync.Map{}
upLevelNode []*node //上游节点
upNodeWrite = make(chan []byte, 999)
extNodeIp []string
@@ -64,7 +86,7 @@ func InitCurrentNode() {
addr: currentNode.addr,
}
currentNode.mirrorNode.mirrorNode = currentNode
nodeMap[currentNode.uuid] = currentNode
nodeMap.Store(currentNode.uuid, currentNode)
//fmt.Println("当前节点UUID", currentNode.uuid)
go func() {
for b := range upNodeWrite {
@@ -91,6 +113,40 @@ func InitCurrentNode() {
nodeTickPing()
time.AfterFunc(time.Second*10, checkUpLevelNode)
}
func getNode(arg string) (n *node, err error) {
id, err := strconv.Atoi(arg)
if err == nil {
nodeMap.Range(func(key, value interface{}) bool {
_n := value.(*node)
if _n.id == id {
n = _n
return false
}
return true
})
} else {
nodeMap.Range(func(key, value interface{}) bool {
node := value.(*node)
if fmt.Sprintf("%s:%d", node.mainIp, node.port) == arg {
n = node
return false
} else if fmt.Sprintf("%s:%d", node.addr, node.port) == arg {
n = node
return false
} else if node.uuid == arg {
n = node
return false
}
return true
})
}
if n != nil {
return n, nil
} else {
return connectNew(arg)
}
}
func checkUpLevelNode() {
if len(currentConfig.DstNode) > 0 && len(upLevelNode) == 0 {
@@ -114,36 +170,20 @@ func checkUpLevelNode() {
return
}
}
func() {
l := clientLock.RLock()
defer l.RUnlock()
for _, n := range nodeMap {
if n.uuid != currentNode.uuid {
func() {
l.RUnlock()
defer clientLock.RLock(l)
if len(n.mainIp) == 0 {
_, err := getNode(fmt.Sprintf("%s:%d", n.addr, n.port))
if common.Debug {
fmt.Printf("连接n.addr %s 错误 %v \r\n", fmt.Sprintf("%s:%d", n.addr, n.port), err)
}
}
}()
if len(upLevelNode) > 0 {
return
}
nodeMap.Range(func(key, value interface{}) bool {
n := value.(*node)
if n.uuid != currentNode.uuid {
if len(n.mainIp) == 0 {
getNode(fmt.Sprintf("%s:%d", n.addr, n.port))
}
if len(upLevelNode) > 0 {
return false
}
}
}()
return true
})
}
}
}
time.AfterFunc(time.Second*5, checkUpLevelNode)
}
@@ -153,7 +193,8 @@ func nodeTickPing() {
defer l.RUnlock()
now := time.Now().Unix()
for _, n := range nodeMap {
nodeMap.Range(func(key, value interface{}) bool {
n := value.(*node)
if n.uuid != currentNode.uuid {
if n.mainIp != "" {
addr1 := fmt.Sprintf("%s:%d", n.mainIp, n.port)
@@ -177,8 +218,8 @@ func nodeTickPing() {
}
}
}
return true
})
time.AfterFunc(time.Second*1, nodeTickPing)
}
@@ -288,19 +329,19 @@ func connectNew(addr string) (n *node, e error) {
}
n.mainIp = cert.RSADecrypterStr(regmsg.MainIp)
if n.port, err = strconv.Atoi(cert.RSADecrypterStr(regmsg.Port)); err != nil {
if n.port, err = strconv.Atoi(cert.RSADecrypterStr(regmsg.Port)); n.port==0 {
n.port = -1
}
if v, ok := nodeMap[regmsg.UUID]; ok {
if v.conn.node != nil && v.conn.node.uuid == regmsg.UUID && v.conn.closeTag == 0 {
if v, ok := nodeMap.Load(regmsg.UUID); ok {
if v.(*node).conn.node != nil && v.(*node).conn.node.uuid == regmsg.UUID && v.(*node).conn.closeTag == 0 {
n.uuid = "" //清空uuid避免正常的node被删
n.conn.Close("重复注册") //当前的连接关掉
n.conn = v.conn
v.mainIp = cert.RSADecrypterStr(regmsg.MainIp)
if v.port, err = strconv.Atoi(cert.RSADecrypterStr(regmsg.Port)); err != nil {
v.port = -1
v.(*node).mainIp = cert.RSADecrypterStr(regmsg.MainIp)
if v.(*node).port, err = strconv.Atoi(cert.RSADecrypterStr(regmsg.Port)); v.(*node).port==0 {
v.(*node).port = -1
}
n = v
n = v.(*node)
} else {
n.conn.node = n
}
@@ -308,10 +349,9 @@ func connectNew(addr string) (n *node, e error) {
} else {
n.conn.node = n
}
nodeMap[n.uuid] = n
nodeMap.Store(n.uuid, n)
n.reConnectAddrs = []string{addr}
n.Write(common.CMD_GET_NODE, 0, nil)
return n, nil
case <-time.After(time.Second * 10):
return nil, errors.New("time out")
@@ -454,9 +494,6 @@ func (n *node) do(msg *common.Msg) {
}
case common.CMD_REG:
func() {
l := clientLock.Lock()
defer l.Unlock()
var regmsg RegMsg
err = json.Unmarshal(msg.CmdData, &regmsg)
@@ -468,7 +505,7 @@ func (n *node) do(msg *common.Msg) {
}
uuid := regmsg.UUID
if uuid == currentNode.uuid {
regmsg.Err = "不能连接自己"
regmsg.Err = "请求的UUID相同,无法连接自己,请将节点设置为不同的UUID"
b, _ := json.Marshal(regmsg)
n.Write(common.CMD_REG_RESULT, 0, b)
return
@@ -491,23 +528,20 @@ func (n *node) do(msg *common.Msg) {
resultMsg.Port = cert.RSAEncrypterStr(strconv.Itoa(currentNode.port))
resultMsg.Goos = cert.RSAEncrypterStr(currentNode.goos)
b, _ := json.Marshal(resultMsg)
//返回成功结果
n.Write(common.CMD_REG_RESULT, 0, b)
//储存节点
n.uuid = uuid
if v, ok := nodeMap[uuid]; !ok || v.conn.closeTag > 0 {
n.conn.node = n
if common.Debug {
fmt.Printf("nodeMap2 %s %p \r\n", regmsg.UUID, n)
}
nodeMap[regmsg.UUID] = n
b, _ := json.Marshal(resultMsg)
//返回成功结果
n.Write(common.CMD_REG_RESULT, 0, b)
//储存节点
n.uuid = uuid
if v, ok := nodeMap.Load(uuid); !ok || v.(*node).conn.closeTag > 0 {
n.conn.node = n
if common.Debug {
fmt.Printf("nodeMap2 %s %p \r\n", regmsg.UUID, n)
}
currentNode.broadcastNode()
//把本机所有节点同步到注册机器
go n.writeGetNodeResult(msg.CmdId)
}()
nodeMap.Store(regmsg.UUID, n)
}
currentNode.broadcastNode()
case common.CMD_REG_RESULT:
var regmsg RegMsg
err = json.Unmarshal(msg.CmdData, &regmsg)
@@ -520,7 +554,7 @@ func (n *node) do(msg *common.Msg) {
default:
}
//交换节点
go n.writeGetNodeResult(msg.CmdId)
n.writeGetNodeResult(msg.CmdId)
case common.CMD_REMOTE_REG:
var regmsg RegMsg
@@ -569,9 +603,11 @@ func (n *node) do(msg *common.Msg) {
return
}
l := clientLock.Lock()
defer l.Unlock()
if n.uuid != regmsg.UUID {
var targetNode *node
if targetNode, ok = nodeMap[regmsg.UUID]; !ok {
if _v, ok := nodeMap.Load(regmsg.UUID); !ok {
targetNode = getNewNode(nodeInfo{
UUID: regmsg.UUID,
HostName: cert.RSADecrypterStr(regmsg.Hostname),
@@ -582,8 +618,9 @@ func (n *node) do(msg *common.Msg) {
if common.Debug {
fmt.Printf("nodeMap4 %s %p \r\n", regmsg.UUID, n)
}
nodeMap[regmsg.UUID] = targetNode
nodeMap.Store(regmsg.UUID, targetNode)
} else {
targetNode = _v.(*node)
targetNode.updateNode(nodeInfo{
UUID: regmsg.UUID,
HostName: cert.RSADecrypterStr(regmsg.Hostname),
@@ -596,8 +633,6 @@ func (n *node) do(msg *common.Msg) {
} else {
v <- n
}
l.Unlock()
n.writeGetNodeResult(msg.CmdId)
case common.CMD_PING:
@@ -805,12 +840,12 @@ func (n *node) do(msg *common.Msg) {
Goos: cert.RSADecrypterStr(_n.Goos),
}
if _n.UUID != currentNode.uuid {
if v, ok := nodeMap[_n.UUID]; !ok {
nodeMap[_n.UUID] = getNewNode(_n, n)
if v, ok := nodeMap.Load(_n.UUID); !ok {
nodeMap.Store(_n.UUID, getNewNode(_n, n))
} else {
v.hostName = _n.HostName
v.mainIp = _n.MainIp
v.port, _ = strconv.Atoi(_n.Port)
v.(*node).hostName = _n.HostName
v.(*node).mainIp = _n.MainIp
v.(*node).port, _ = strconv.Atoi(_n.Port)
}
}
@@ -842,34 +877,29 @@ func (n *node) do(msg *common.Msg) {
if err != nil {
return
}
l := clientLock.Lock()
defer l.Unlock()
if v, ok := nodeMap[nmsg.UUID]; !ok {
if v, ok := nodeMap.Load(nmsg.UUID); !ok {
newNode := getNewNode(nmsg, n)
if common.Debug {
fmt.Printf("nodeMap5 %s %p \r\n", nmsg.UUID, newNode)
}
nodeMap[nmsg.UUID] = newNode
nodeMap.Store(nmsg.UUID, newNode)
} else if nmsg.UUID != currentNode.uuid {
n := v.(*node)
port, err := strconv.Atoi(cert.RSADecrypterStr(nmsg.Port))
if err == nil {
v.port = port
n.port = port
} else {
v.port = -1
n.port = -1
}
v.mainIp = cert.RSADecrypterStr(nmsg.MainIp)
v.hostName = cert.RSADecrypterStr(nmsg.HostName)
v.goos = cert.RSADecrypterStr(nmsg.Goos)
v.uuid = nmsg.UUID
n.mainIp = cert.RSADecrypterStr(nmsg.MainIp)
n.hostName = cert.RSADecrypterStr(nmsg.HostName)
n.goos = cert.RSADecrypterStr(nmsg.Goos)
n.uuid = nmsg.UUID
if common.Debug {
fmt.Printf("nodeMap6 %s %p \r\n", nmsg.UUID, v)
}
nodeMap[nmsg.UUID] = v
nodeMap.Store(nmsg.UUID, n)
}
case common.CMD_DIR:
@@ -1240,26 +1270,17 @@ func getNewNode(m nodeInfo, n *node) *node {
func allNodesDo(f func(*node) (bool, error)) (err error) {
var ok bool
l := clientLock.RLock()
defer l.RUnlock()
for _, n := range nodeMap {
nodeMap.Range(func(key, value interface{}) bool {
n := value.(*node)
if n.uuid != currentNode.uuid {
func() {
l.RUnlock()
defer clientLock.RLock(l)
ok, err = f(n)
}()
if err != nil {
return err
}
if !ok {
break
ok, err = f(n)
if err != nil || !ok {
return false
}
}
}
return nil
return true
})
return err
}
func (n *node) ping(id uint32) {
if common.NoPing {
@@ -1332,12 +1353,7 @@ func (n *node) ping(id uint32) {
func (n *node) Delete(reason string) {
go func() {
if atomic.CompareAndSwapInt32(&n.isClose, 0, 1) {
l := clientLock.Lock()
_, ok := nodeMap[n.uuid]
if ok {
delete(nodeMap, n.uuid)
}
l.Unlock()
n.connMap.Range(func(key, value interface{}) bool {
if v, ok := value.(common.Conn); ok {
v.Close(reason)
@@ -1368,6 +1384,7 @@ func (n *node) Delete(reason string) {
n.shellMap.Delete(key)
return true
})
nodeMap.Delete(n.uuid)
}
}()
@@ -1464,27 +1481,32 @@ func (n *node) storeConn(v common.Conn) (newID uint32) {
}
func (n *node) writeGetNodeResult(id uint32) {
l := clientLock.RLock()
go func() {
defer func() {
if err := recover(); err != nil {
fmt.Println(err)
debug.PrintStack()
}
}()
var s []*nodeInfo
nodeMap.Range(func(key, value interface{}) bool {
_n := value.(*node)
if _n.uuid != currentNode.uuid {
s = append(s, &nodeInfo{
UUID: _n.uuid,
HostName: cert.RSAEncrypterStr(_n.hostName),
MainIp: cert.RSAEncrypterStr(_n.mainIp),
Port: cert.RSAEncrypterStr(strconv.Itoa(_n.port)),
Goos: cert.RSAEncrypterStr(_n.goos),
})
}
return true
})
defer l.RUnlock()
b, _ := json.Marshal(s)
n.Write(common.CMD_GET_NODE_RESULT, id, b)
}()
var s []*nodeInfo
for _, _n := range nodeMap {
if _n.uuid != currentNode.uuid {
s = append(s, &nodeInfo{
UUID: _n.uuid,
HostName: cert.RSAEncrypterStr(_n.hostName),
MainIp: cert.RSAEncrypterStr(_n.mainIp),
Port: cert.RSAEncrypterStr(strconv.Itoa(_n.port)),
Goos: cert.RSAEncrypterStr(_n.goos),
})
}
}
b, _ := json.Marshal(s)
n.Write(common.CMD_GET_NODE_RESULT, id, b)
}
func (n *node) updateNode(msg nodeInfo) {
n.hostName = msg.HostName
-832
View File
@@ -1,832 +0,0 @@
package server
/*
*高级shell功能
*node节点管理、remoteShell远程shellconfig配置管理
*/
import (
"bytes"
"cert"
"encoding/json"
"fmt"
"io"
"io/ioutil"
"os"
"os/exec"
"rakshasa/aes"
"rakshasa/common"
"github.com/google/uuid"
"regexp"
"runtime"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/abiosoft/readline"
"github.com/creack/pty"
"github.com/luyu6056/ishell"
"golang.org/x/text/encoding/simplifiedchinese"
"golang.org/x/text/transform"
)
var (
shellMapLock sync.Mutex
)
type StartCmdParam struct {
Param string
Size *pty.Winsize
}
type remoteCmd struct {
cmdStatus int32
cmd *exec.Cmd
id uint32
stdin io.WriteCloser
inChan chan []byte
translate func(in []byte) ([]byte, error)
ping, pong int64
}
func init() {
configShell := cliInit()
configShell.SetPrompt("rakshasa\\config>")
configShell.AddCmd(&ishell.Cmd{
Name: "info",
Help: "打印当前配置",
Func: func(c *ishell.Context) {
c.Println("当前节点", currentNode.uuid)
c.Println("上级节点地址", currentConfig.DstNode)
c.Println("通讯密码", currentConfig.Password)
c.Println("监听端口", currentConfig.Port)
c.Println("监听IP", currentConfig.ListenIp)
c.Println("禁止额外连接", currentConfig.Limit)
c.Println("配置文件名", currentConfig.FileName)
if currentConfig.FileSave {
c.Println("当前配置:已写入文件")
} else {
c.Println("当前配置:未写入文件")
}
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "save",
Help: "保存文件",
Func: func(c *ishell.Context) {
if err := ConfigSave(); err == nil {
c.Println("写入成功")
} else {
c.Println("保存失败", err.Error())
}
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "d",
Help: "修改上级节点地址,格式为 ip:端口 多个节点以,隔开 注意:不会立刻连接设置节点, 当发生 节点掉线重连 时候会连接该地址",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,格式为 ip:端口 多个节点以,隔开 如 d 192.168.1.1:8883,192.168.1.2:8883")
return
}
dstNode, err := common.ResolveTCPAddr(c.Args[0])
if err != nil {
c.Println("参数错误,格式为 ip:端口 多个节点以,隔开 如 d 192.168.1.1:8883,192.168.1.2:8883")
return
}
currentConfig.DstNode = dstNode
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "password",
Help: "修改通讯密码,立即生效",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,格式为 password \"123456\"")
return
}
c.Println(c.Args)
currentConfig.Password = c.Args[0]
currentConfig.FileSave = false
aes.Key = aes.MD5_B(currentConfig.Password + string(cert.RsaPrivateKey[:16]))
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "port",
Help: "修改监听端口,立即生效",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,格式为 port 8883")
return
}
port, _ := strconv.Atoi(c.Args[0])
if port <= 0 || port > 65535 {
c.Println("参数错误,端口范围是1-65535")
return
}
c.Println("正在关闭server监听")
if currentNode.listen != nil {
currentNode.listen.Close()
currentNode.listen = nil
}
currentConfig.Port = port
currentNode.port = port
currentConfig.FileSave = false
if err := StartServer(fmt.Sprintf(":%d", currentConfig.Port)); err != nil {
c.Printf("启动节点失败 %v, 请重新修改监听端口", currentConfig.Port)
}
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "ip",
Help: "修改本节点连接ip,当其他节点进行额外连接时候,优先使用此ip连接",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
currentConfig.ListenIp = c.Args[0]
currentNode.mainIp = currentConfig.ListenIp
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "limit",
Help: "修改本节点Limit设置,使用方法 limit true",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
currentConfig.Limit = c.Args[0] == "true"
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "f",
Help: "修改配置文件名,使用方法 f config.yaml",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
currentConfig.FileName = c.Args[0]
currentConfig.FileSave = false
},
})
configShell.AddCmd(&ishell.Cmd{
Name: "uuid",
Help: "修改本节点UUID设置,使用方法uuid 字串符",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
if id, err := uuid.Parse(c.Args[0]); err == nil {
currentConfig.UUID = id.String()
currentConfig.FileSave = false
SetConfig(currentConfig)
} else {
c.Println("输入的uuid不是合法的uuid,建议使用xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx")
}
},
})
rootCli.AddCmd(&ishell.Cmd{
Name: "config",
Help: "配置管理",
Func: func(c *ishell.Context) {
configShell.Run()
},
})
remoteShell := cliInit()
remoteShell.SetPrompt("rakshasa\\remoteshell>")
fileShell := cliInit()
remoteShell.AddCmd(&ishell.Cmd{
Name: "file",
Help: "连到节点进行文件管理,参数为id或者uuid",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
workN, _ := getNode(c.Args[0])
if workN == nil {
c.Println("无法连接节点", c.Args[0])
return
}
if workN != nil {
fileShell.Set("node", workN)
result := make(chan interface{}, 1)
id := workN.storeQuery(result)
workN.Write(common.CMD_PWD, id, []byte(cert.RSAEncrypterByPriv(currentConfig.Password)))
select {
case pwd := <-result:
workN.deleteQuery(id)
pwd = strings.ReplaceAll(pwd.(string), "\\", "/")
fileShell.Set("pwd", pwd)
fileShell.SetPrompt(workN.uuid + " " + pwd.(string) + ">")
fileShell.Run()
case <-time.After(common.CMD_TIMEOUT):
workN.deleteQuery(id)
c.Println("连接", c.Args[0], "超时")
}
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "dir",
Help: "打印当前目录文件",
Func: func(c *ishell.Context) {
pwd := fileShell.Get("pwd")
n := c.Get("node").(*node)
resChan := make(chan interface{}, 1)
id := n.storeQuery(resChan)
n.Write(common.CMD_DIR, id, []byte(cert.RSAEncrypterByPriv(pwd.(string))))
select {
case res := <-resChan:
n.deleteQuery(id)
c.Println(res)
case <-time.After(common.CMD_TIMEOUT):
n.deleteQuery(id)
c.Println("dir time out")
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "cd",
Help: "切换工作目录",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误")
return
}
dir := c.Args[0]
pwd := fileShell.Get("pwd").(string)
n := c.Get("node").(*node)
if strings.Contains(dir, ":/") || dir[0] == '/' || dir == "~" {
pwd = dir
} else {
pwd += "/" + dir
pwd = strings.TrimRight(getRealPath(pwd), "/")
}
resChan := make(chan interface{}, 1)
id := n.storeQuery(resChan)
n.Write(common.CMD_CD, id, []byte(cert.RSAEncrypterByPriv(pwd)))
select {
case res := <-resChan:
n.deleteQuery(id)
if err, ok := res.(error); ok {
c.Println(err.Error())
} else {
pwd = res.(string)
fileShell.Set("pwd", pwd)
c.SetPrompt(n.uuid + " " + pwd + ">")
}
case <-time.After(common.CMD_TIMEOUT):
n.deleteQuery(id)
c.Println("dir time out")
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "upload",
Help: "上传文件 ,upload 本地文件 远程目录(为空传到工作目录)",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 && len(c.Args) != 2 {
c.Println("参数错误")
return
}
s, err := os.Stat(c.Args[0])
if err != nil {
c.Println("打开本地文件", c.Args[0], "错误 ", err)
return
}
f, err := os.Open(c.Args[0])
if err != nil {
c.Println("打开本地文件", c.Args[0], "错误 ", err)
return
}
defer f.Close()
pwd := fileShell.Get("pwd").(string) + "/"
n := c.Get("node").(*node)
if len(c.Args) == 2 {
pwd = c.Args[1]
}
pwd = strings.ReplaceAll(pwd, "\\", "/")
c.Args[0] = strings.ReplaceAll(c.Args[0], "\\", "/")
i := strings.LastIndex(c.Args[0], "/")
if i == -1 {
i = 0
}
if pwd[len(pwd)-1] == '/' {
pwd += c.Args[0][i:]
}
i = strings.LastIndex(pwd, "/")
if i == -1 {
i = 0
}
filename := pwd[i+1:]
dir := pwd[:i]
dir = strings.TrimRight(getRealPath(dir), "/") + "/"
pwd = dir + filename
resChan := make(chan interface{}, 9999) //避免收消息阻塞
filereadChan := make(chan []byte, 10)
upload := func() {
for i := 0; i < 10; i++ {
buf := make([]byte, common.MAX_PACKAGE-len(pwd)-9)
n, err := f.Read(buf)
if err != nil {
if err == io.EOF {
return
}
resChan <- err
c.Println("读取文件", c.Args[0], "错误", err)
return
}
filereadChan <- buf[:n]
}
}
offset := 0
be := len(pwd) + 1
id := n.storeQuery(resChan)
defer n.deleteQuery(id)
b := []byte(pwd)
b = append(b, 0, 0, 0, 0, 0, 0, 0, 0, 0)
c.ProgressBar().Start()
go upload()
var resnum int
for {
select {
case data := <-filereadChan:
b[be] = byte(offset)
b[be+1] = byte(offset >> 8)
b[be+2] = byte(offset >> 16)
b[be+3] = byte(offset >> 24)
b[be+4] = byte(offset >> 32)
b[be+5] = byte(offset >> 40)
b[be+6] = byte(offset >> 48)
b[be+7] = byte(offset >> 56)
offset += len(data)
n.Write(common.CMD_UPLOAD, id, cert.RSAEncrypterByPrivByte(append(b, data...)))
case res := <-resChan:
switch v := res.(type) {
case error:
c.ProgressBar().Stop()
c.Println("上传失败", res)
return
case int64:
resnum++
i := v * 100 / s.Size()
c.ProgressBar().Suffix(fmt.Sprint(" ", i, "%"))
c.ProgressBar().Progress(int(i))
if v == s.Size() {
c.ProgressBar().Stop()
c.Println(c.Args[0], "上传成功")
return
}
if resnum >= 5 {
go upload()
resnum -= 10
}
default:
c.Println("协议错误")
return
}
case <-time.After(common.CMD_TIMEOUT):
c.ProgressBar().Stop()
c.Println("upload time out")
return
}
}
},
})
fileShell.AddCmd(&ishell.Cmd{
Name: "download",
Help: "下载文件 download 远程文件 本地目录(为空本地执行目录)",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 && len(c.Args) != 2 {
c.Println("参数错误")
return
}
pwd := fileShell.Get("pwd").(string)
n := c.Get("node").(*node)
file := c.Args[0]
file = strings.ReplaceAll(file, "\\", "/")
if strings.Contains(file, ":/") || file[0] == '/' {
pwd = file
} else {
pwd += "/" + file
}
i := strings.LastIndex(pwd, "/")
if i == -1 {
i = 0
}
filename := pwd[i+1:]
dir := pwd[:i]
dir = strings.TrimRight(getRealPath(dir), "/") + "/"
mydir, err := os.Getwd()
local := "./" + filename
if err == nil {
local = mydir + "/" + filename
}
if len(c.Args) == 2 {
s, err := os.Stat(c.Args[1])
if err == nil {
if s.IsDir() {
local = strings.TrimRight(c.Args[1], "/") + "/" + filename
} else {
local = c.Args[1]
}
} else {
local = c.Args[1]
}
}
pwd = dir + filename
result := make(chan interface{}, 999)
id := n.storeQuery(result)
defer n.deleteQuery(id)
b := []byte(pwd)
b = append(b, []byte{0, 0, 0, 0, 0, 0, 0, 0, 0}...)
total := int64(-1)
be := len(pwd) + 1
b[be] = byte(total)
b[be+1] = byte(total >> 8)
b[be+2] = byte(total >> 16)
b[be+3] = byte(total >> 24)
b[be+4] = byte(total >> 32)
b[be+5] = byte(total >> 40)
b[be+6] = byte(total >> 48)
b[be+7] = byte(total >> 56)
n.Write(common.CMD_DOWNLOAD, id, cert.RSAEncrypterByPrivByte(b))
c.ProgressBar().Start()
size := int64(0)
resnum := 0
total = 0
var f *os.File
for {
select {
case res := <-result:
switch v := res.(type) {
case error:
c.ProgressBar().Stop()
c.Println("下载失败", res)
return
case int64:
var err error
size = v
f, err = os.OpenFile(local, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0666)
if err != nil {
c.Println("本地文件 ", local, "写入失败", err.Error())
return
}
defer f.Close()
case []byte:
if f == nil {
c.Println("本地文件 ", local, "不可写入")
return
}
resnum++
num, err := f.Write(v)
if err != nil {
c.Println("本地文件 ", local, "写入失败", err.Error())
return
}
if num != len(v) {
c.Println("本地文件 ", local, "写入失败,写入量不符")
return
}
total += int64(num)
i := total * 100 / size
c.ProgressBar().Suffix(fmt.Sprint(" ", i, "%"))
c.ProgressBar().Progress(int(i))
if total == size {
c.ProgressBar().Stop()
c.Println(c.Args[0], "下载成功 文件保存到", local)
return
}
if resnum == 10 {
resnum -= 10
b[be] = byte(total)
b[be+1] = byte(total >> 8)
b[be+2] = byte(total >> 16)
b[be+3] = byte(total >> 24)
b[be+4] = byte(total >> 32)
b[be+5] = byte(total >> 40)
b[be+6] = byte(total >> 48)
b[be+7] = byte(total >> 56)
n.Write(common.CMD_DOWNLOAD, id, cert.RSAEncrypterByPrivByte(b))
}
default:
c.Println("协议错误")
return
}
case <-time.After(common.CMD_TIMEOUT):
c.ProgressBar().Stop()
c.Println("upload time out")
return
}
}
},
})
remoteShell.AddCmd(&ishell.Cmd{
Name: "new",
Help: "与一个或者多个节点连接,使用方法 new ip:端口 多个地址以,间隔 如1080 127.0.0.1:1081,127.0.0.1:1082",
Func: func(c *ishell.Context) {
if len(c.Args) != 1 {
c.Println("参数错误,使用方法 connect ip:端口")
return
}
for _, addr := range strings.Split(c.Args[0], ",") {
_, err := getNode(addr)
if err != nil {
c.Println("连接", addr, "失败", err)
return
}
}
},
})
remoteShell.AddCmd(&ishell.Cmd{
Name: "shell",
Help: "反弹shell 使用方法 shell id/uuid 启动参数 ,启动参数可为空,win默认启动cmd,linux默认启动bash 如 shell 1 powershell 。 shell 1 zsh",
Func: func(c *ishell.Context) {
if len(c.Args) < 1 {
c.Println("参数错误,例子 shell 1 powershell")
return
}
param := ""
if len(c.Args) == 2 {
param = c.Args[1]
}
n, _ := getNode(c.Args[0])
if n == nil {
c.Println("无法连接节点", c.Args[0])
return
}
res := make(chan interface{}, 999)
id := n.storeQuery(res)
defer n.deleteQuery(id)
p := StartCmdParam{
Param: param,
Size: common.GetSize(),
}
b, _ := json.Marshal(p)
n.Write(common.CMD_SHELL, id, cert.RSAEncrypterByPrivByte(b))
s := &remoteCmd{
cmd: nil,
stdin: nil,
inChan: make(chan []byte, 999),
translate: func(in []byte) ([]byte, error) { return in, nil },
pong: time.Now().Unix(),
}
select {
case i := <-res:
switch v := i.(type) {
case error:
c.Println("启动shell失败,错误", v.Error())
case []byte:
data := v
s.id = uint32(data[0]) | uint32(data[1])<<8 | uint32(data[2])<<16 | uint32(data[3])<<24
switch data[4] {
case 0: //windows
if string(data[len(data)-6:]) == string([]byte{32, 57, 51, 54, 13, 10}) { //活动代码页: 936
//gbk转utf8
s.translate = func(in []byte) ([]byte, error) {
reader := transform.NewReader(bytes.NewReader(in), simplifiedchinese.GBK.NewDecoder())
d, e := ioutil.ReadAll(reader)
if e != nil {
return nil, e
}
return d, nil
}
}
case 1: //linux
if runtime.GOOS == "windows" {
if !common.EnableTermVt {
s.translate = func(in []byte) ([]byte, error) {
if in[0] == 27 {
r, _ := regexp.Compile(`\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])`)
res := r.ReplaceAllString(string(in), "")
return []byte(res), nil
}
return in, nil
}
}
}
}
atomic.CompareAndSwapInt32(&s.cmdStatus, 0, 1)
}
case <-time.After(common.CMD_TIMEOUT):
c.Println("启动shell失败,超时")
return
}
n.shellMap.Store(s.id, s)
r, _ := readline.NewEx(&readline.Config{FuncIsTerminal: func() bool { return false }, ForcePrint: true})
defer func() {
n.shellMap.Delete(s.id)
atomic.StoreInt32(&s.cmdStatus, -1)
c.Println("请按回车键退出")
r.Close()
}()
go func() {
for {
switch s.cmdStatus {
case 1:
input, err := r.ReadlineEx()
if err != nil {
if err != readline.ErrInterrupt {
res <- err
return
}
if s.cmdStatus == 1 {
n.Write(common.CMD_SHELL_DATA, s.id, []byte{03})
}
}
if s.cmdStatus == 1 {
n.Write(common.CMD_SHELL_DATA, s.id, []byte(input+"\n"))
}
case 0:
time.Sleep(time.Millisecond * 100)
case -1:
return
}
}
}()
tick := time.NewTicker(common.CMD_TIMEOUT / 2)
for {
select {
case b := <-s.inChan:
s.pong = time.Now().Unix()
if len(b) > 0 {
b, err := s.translate(b)
if err != nil {
c.Println("shell 运行失败", err)
return
}
fmt.Print(string(b))
}
case v := <-res:
if err, ok := v.(error); ok {
if err.Error() != "退出shell" {
c.Println("运行shell", param, "失败", err)
}
} else {
c.Println("无法处理消息", v)
}
return
case <-tick.C:
s.ping = time.Now().Unix()
if s.ping-s.pong > int64(common.CMD_TIMEOUT/time.Second) {
c.Println("shell time out")
return
}
n.Write(common.CMD_SHELL_DATA, s.id, nil)
}
}
},
})
rootCli.AddCmd(&ishell.Cmd{
Name: "remoteshell",
Help: "远程shell",
Func: func(c *ishell.Context) {
remoteShell.Run()
},
})
}
// 打印节点
func printNodes(c *ishell.Context) {
l := clientLock.RLock()
defer l.RUnlock()
var list []*node
for _, n := range nodeMap {
list = append(list, n)
}
orderNode(list)
c.Println("ID UUID HostName GOOS IP listenIP")
c.Println("-----------------------------------------------------------------------------------------------------------------------------")
for k, n := range list {
n.id = k + 1
hostname := bytes.Repeat([]byte(" "), 22)
copy(hostname, n.hostName)
ip := bytes.Repeat([]byte(" "), 23)
if n.uuid == currentNode.uuid {
copy(ip, "(localhost)"+":"+strconv.Itoa(n.port))
} else {
copy(ip, n.addr+":"+strconv.Itoa(n.port))
}
listenip := n.mainIp
goos := bytes.Repeat([]byte(" "), 11)
copy(goos, n.goos)
c.Printf("%2d %s %s %s %s %s\n", n.id, n.uuid, hostname, goos, ip, listenip)
}
}
func getRealPath(path string) string {
path_s := strings.Split(path, "/")
realpath := []string{}
if len(path_s) == 0 {
return "error"
}
for _, value := range path_s {
if value == ".." {
k := len(realpath)
kk := k - 1
realpath = append(realpath[:kk], realpath[k:]...)
} else {
realpath = append(realpath, value)
}
}
return strings.Join(realpath, "/")
}
func printConn() {
connMap.Range(func(key, value interface{}) bool {
fmt.Println(key)
return true
})
}
func getNode(arg string) (*node, error) {
l := clientLock.Lock()
defer l.Unlock()
id, err := strconv.Atoi(arg)
if err == nil {
for _, n := range nodeMap {
if n.id == id {
return n, nil
}
}
} else {
for _, node := range nodeMap {
if fmt.Sprintf("%s:%d", node.mainIp, node.port) == arg {
return node, nil
} else if fmt.Sprintf("%s:%d", node.addr, node.port) == arg {
return node, nil
} else if node.uuid == arg {
return node, nil
}
}
}
return connectNew(arg)
}