feat: add file transfer support (#55)

* feat: add file transfer support

* 1MB buffer
This commit is contained in:
UUBulb
2024-08-20 22:24:03 +08:00
committed by GitHub
parent 093275bc80
commit 73a727d435
9 changed files with 342 additions and 54 deletions
+65
View File
@@ -0,0 +1,65 @@
package fm
import (
"bytes"
"encoding/binary"
)
var (
fileIdentifier = []byte{0x4E, 0x5A, 0x54, 0x44} // NZTD
fileNameIdentifier = []byte{0x4E, 0x5A, 0x46, 0x4E} // NZFN
errorIdentifier = []byte{0x4E, 0x45, 0x52, 0x52} // NERR
completeIdentifier = []byte{0x4E, 0x5A, 0x55, 0x50} // NZUP
)
func AppendFileName(bin []byte, data string, isDir bool) []byte {
buffer := bytes.NewBuffer(bin)
appendFileName(buffer, isDir, []byte(data))
return buffer.Bytes()
}
func Create(buffer *bytes.Buffer, path string) []byte {
// Write identifier for TypeFileName (4 bytes)
binary.Write(buffer, binary.BigEndian, fileNameIdentifier)
// Write length of path (4 byte)
binary.Write(buffer, binary.BigEndian, uint32(len(path)))
// Write path string
binary.Write(buffer, binary.BigEndian, []byte(path))
return buffer.Bytes()
}
func CreateFile(buffer *bytes.Buffer, size uint64) []byte {
// Write identifier for TypeFile (4 bytes)
binary.Write(buffer, binary.BigEndian, fileIdentifier)
// Write file size (8 bytes)
binary.Write(buffer, binary.BigEndian, size)
return buffer.Bytes()
}
func CreateErr(err error) []byte {
buffer := new(bytes.Buffer)
binary.Write(buffer, binary.BigEndian, errorIdentifier)
binary.Write(buffer, binary.BigEndian, []byte(err.Error()))
return buffer.Bytes()
}
func appendFileName(buffer *bytes.Buffer, isDir bool, data []byte) {
// Write file type (1 byte)
if isDir {
binary.Write(buffer, binary.BigEndian, byte(1))
} else {
binary.Write(buffer, binary.BigEndian, byte(0))
}
// Write the length of file name (1 byte)
length := byte(len(data))
binary.Write(buffer, binary.BigEndian, length)
// Write file name
buffer.Write(data)
}
+158
View File
@@ -0,0 +1,158 @@
package fm
import (
"bytes"
"encoding/binary"
"errors"
"io"
"io/fs"
"os"
"os/user"
"path/filepath"
pb "github.com/nezhahq/agent/proto"
)
type Task struct {
taskClient pb.NezhaService_IOStreamClient
printf func(string, ...interface{})
remoteData *pb.IOStreamData
}
func NewFMClient(client pb.NezhaService_IOStreamClient, printFunc func(string, ...interface{})) *Task {
return &Task{
taskClient: client,
printf: printFunc,
}
}
func (t *Task) DoTask(data *pb.IOStreamData) {
t.remoteData = data
switch t.remoteData.Data[0] {
case 0:
t.listDir()
case 1:
go t.download()
case 2:
t.upload()
}
}
func (t *Task) listDir() {
dir := string(t.remoteData.Data[1:])
var entries []fs.DirEntry
var err error
for {
entries, err = os.ReadDir(dir)
if err != nil {
usr, err := user.Current()
if err != nil {
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
dir = usr.HomeDir + string(filepath.Separator)
continue
}
break
}
var buffer bytes.Buffer
td := Create(&buffer, dir)
for _, e := range entries {
newBin := AppendFileName(td, e.Name(), e.IsDir())
td = newBin
}
t.taskClient.Send(&pb.IOStreamData{Data: td})
}
func (t *Task) download() {
path := string(t.remoteData.Data[1:])
file, err := os.Open(path)
if err != nil {
println("Error opening file: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
defer file.Close()
fileInfo, err := file.Stat()
if err != nil {
println("Error getting file info: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
fileSize := fileInfo.Size()
if fileSize <= 0 {
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(errors.New("requested file is empty"))})
return
}
// Send header (12 bytes)
var header bytes.Buffer
headerData := CreateFile(&header, uint64(fileSize))
if err := t.taskClient.Send(&pb.IOStreamData{Data: headerData}); err != nil {
println("Error sending file header: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
buffer := make([]byte, 1048576)
for {
n, err := file.Read(buffer)
if err != nil {
if err == io.EOF {
return
}
println("Error reading file: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
if err := t.taskClient.Send(&pb.IOStreamData{Data: buffer[:n]}); err != nil {
println("Error sending file chunk: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
}
}
func (t *Task) upload() {
if len(t.remoteData.Data) < 9 {
println("data is invalid")
return
}
fileSize := binary.BigEndian.Uint64(t.remoteData.Data[1:9])
path := string(t.remoteData.Data[9:])
file, err := os.Create(path)
if err != nil {
println("Error creating file: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
defer file.Close()
totalReceived := uint64(0)
t.printf("receiving file: %s, size: %d", file.Name(), fileSize)
for totalReceived < fileSize {
if t.remoteData, err = t.taskClient.Recv(); err != nil {
println("Error receiving data: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
bytesWritten, err := file.Write(t.remoteData.Data)
if err != nil {
println("Error writing to file: ", err)
t.taskClient.Send(&pb.IOStreamData{Data: CreateErr(err)})
return
}
totalReceived += uint64(bytesWritten)
}
t.printf("received file %s.", file.Name())
t.taskClient.Send(&pb.IOStreamData{Data: completeIdentifier}) // NZUP
}
+14 -14
View File
@@ -79,7 +79,7 @@ func GetHost() *model.Host {
var cpuType string
hi, err := host.Info()
if err != nil {
println("host.Info error: ", err)
printf("host.Info error: %v", err)
} else {
if hi.VirtualizationRole == "guest" {
cpuType = "Virtual"
@@ -99,7 +99,7 @@ func GetHost() *model.Host {
ci, err := cpu.Info()
if err != nil {
hostDataFetchAttempts["CPU"]++
println("cpu.Info error: ", err, ", attempt: ", hostDataFetchAttempts["CPU"])
printf("cpu.Info error: %v, attempt: %d", err, hostDataFetchAttempts["CPU"])
} else {
hostDataFetchAttempts["CPU"] = 0
for i := 0; i < len(ci); i++ {
@@ -120,7 +120,7 @@ func GetHost() *model.Host {
ret.GPU, err = gpu.GetGPUModel()
if err != nil {
hostDataFetchAttempts["GPU"]++
println("gpu.GetGPUModel error: ", err, ", attempt: ", hostDataFetchAttempts["GPU"])
printf("gpu.GetGPUModel error: %v, attempt: %d", err, hostDataFetchAttempts["GPU"])
} else {
hostDataFetchAttempts["GPU"] = 0
}
@@ -131,7 +131,7 @@ func GetHost() *model.Host {
mv, err := mem.VirtualMemory()
if err != nil {
println("mem.VirtualMemory error: ", err)
printf("mem.VirtualMemory error: %v", err)
} else {
ret.MemTotal = mv.Total
if runtime.GOOS != "windows" {
@@ -142,7 +142,7 @@ func GetHost() *model.Host {
if runtime.GOOS == "windows" {
ms, err := mem.SwapMemory()
if err != nil {
println("mem.SwapMemory error: ", err)
printf("mem.SwapMemory error: %v", err)
} else {
ret.SwapTotal = ms.Total
}
@@ -163,7 +163,7 @@ func GetState(skipConnectionCount bool, skipProcsCount bool) *model.HostState {
cp, err := cpu.Percent(0, false)
if err != nil || len(cp) == 0 {
statDataFetchAttempts["CPU"]++
println("cpu.Percent error: ", err, ", attempt: ", statDataFetchAttempts["CPU"])
printf("cpu.Percent error: %v, attempt: %d", err, statDataFetchAttempts["CPU"])
} else {
statDataFetchAttempts["CPU"] = 0
ret.CPU = cp[0]
@@ -172,7 +172,7 @@ func GetState(skipConnectionCount bool, skipProcsCount bool) *model.HostState {
vm, err := mem.VirtualMemory()
if err != nil {
println("mem.VirtualMemory error: ", err)
printf("mem.VirtualMemory error: %v", err)
} else {
ret.MemUsed = vm.Total - vm.Available
if runtime.GOOS != "windows" {
@@ -183,7 +183,7 @@ func GetState(skipConnectionCount bool, skipProcsCount bool) *model.HostState {
// gopsutil 在 Windows 下不能正确取 swap
ms, err := mem.SwapMemory()
if err != nil {
println("mem.SwapMemory error: ", err)
printf("mem.SwapMemory error: %v", err)
} else {
ret.SwapUsed = ms.Used
}
@@ -195,7 +195,7 @@ func GetState(skipConnectionCount bool, skipProcsCount bool) *model.HostState {
loadStat, err := load.Avg()
if err != nil {
statDataFetchAttempts["Load"]++
println("load.Avg error: ", err, ", attempt: ", statDataFetchAttempts["Load"])
printf("load.Avg error: %v, attempt: %d", err, statDataFetchAttempts["Load"])
} else {
statDataFetchAttempts["Load"] = 0
ret.Load1 = loadStat.Load1
@@ -208,7 +208,7 @@ func GetState(skipConnectionCount bool, skipProcsCount bool) *model.HostState {
if !skipProcsCount {
procs, err = process.Pids()
if err != nil {
println("process.Pids error: ", err)
printf("process.Pids error: %v", err)
} else {
ret.ProcessCount = uint64(len(procs))
}
@@ -360,7 +360,7 @@ func updateGPUStat(gpuStat *uint64) {
gs, err := gpustat.GetGPUStat()
if err != nil {
statDataFetchAttempts["GPU"]++
println("gpustat.GetGPUStat error: ", err, ", attempt: ", statDataFetchAttempts["GPU"])
printf("gpustat.GetGPUStat error: %v, attempt: %d", err, statDataFetchAttempts["GPU"])
atomicStoreFloat64(gpuStat, gs)
} else {
statDataFetchAttempts["GPU"] = 0
@@ -379,7 +379,7 @@ func updateTemperatureStat() {
temperatures, err := sensors.SensorsTemperatures()
if err != nil {
statDataFetchAttempts["Temperatures"]++
println("host.SensorsTemperatures error: ", err, ", attempt: ", statDataFetchAttempts["Temperatures"])
printf("host.SensorsTemperatures error: %v, attempt: %d", err, statDataFetchAttempts["Temperatures"])
} else {
statDataFetchAttempts["Temperatures"] = 0
tempStat := []model.SensorTemperature{}
@@ -410,6 +410,6 @@ func atomicStoreFloat64(x *uint64, v float64) {
atomic.StoreUint64(x, math.Float64bits(v))
}
func println(v ...interface{}) {
util.Println(agentConfig.Debug, v...)
func printf(format string, v ...interface{}) {
util.Printf(agentConfig.Debug, format, v...)
}
+2 -1
View File
@@ -19,7 +19,8 @@ type Pty struct {
cmd *exec.Cmd
}
func DownloadDependency() {
func DownloadDependency() error {
return nil
}
func Start() (IPty, error) {
+8 -13
View File
@@ -5,7 +5,6 @@ package pty
import (
"fmt"
"io"
"log"
"net/http"
"os"
"os/exec"
@@ -55,12 +54,11 @@ func VersionCheck() bool {
return false
}
func DownloadDependency() {
func DownloadDependency() error {
if !isWin10 {
executablePath, err := getExecutableFilePath()
if err != nil {
fmt.Println("NEZHA>> wintty 获取文件路径失败", err)
return
return fmt.Errorf("winpty 获取文件路径失败: %v", err)
}
winptyAgentExe := filepath.Join(executablePath, "winpty-agent.exe")
@@ -69,27 +67,23 @@ func DownloadDependency() {
fe, errFe := os.Stat(winptyAgentExe)
fd, errFd := os.Stat(winptyAgentDll)
if errFe == nil && fe.Size() > 300000 && errFd == nil && fd.Size() > 300000 {
return
return fmt.Errorf("winpty 文件完整性检查失败")
}
resp, err := http.Get("https://github.com/rprichard/winpty/releases/download/0.4.3/winpty-0.4.3-msvc2015.zip")
if err != nil {
log.Println("NEZHA>> wintty 下载失败", err)
return
return fmt.Errorf("winpty 下载失败: %v", err)
}
defer resp.Body.Close()
content, err := io.ReadAll(resp.Body)
if err != nil {
log.Println("NEZHA>> wintty 下载失败", err)
return
return fmt.Errorf("winpty 下载失败: %v", err)
}
if err := os.WriteFile("./wintty.zip", content, os.FileMode(0777)); err != nil {
log.Println("NEZHA>> wintty 写入失败", err)
return
return fmt.Errorf("winpty 写入失败: %v", err)
}
if err := unzip.New("./wintty.zip", "./wintty").Extract(); err != nil {
fmt.Println("NEZHA>> wintty 解压失败", err)
return
return fmt.Errorf("winpty 解压失败: %v", err)
}
arch := "x64"
if runtime.GOARCH != "amd64" {
@@ -101,6 +95,7 @@ func DownloadDependency() {
os.RemoveAll("./wintty")
os.RemoveAll("./wintty.zip")
}
return nil
}
func getExecutableFilePath() (string, error) {
+2 -1
View File
@@ -14,7 +14,8 @@ type Pty struct {
tty *conpty.ConPty
}
func DownloadDependency() {
func DownloadDependency() error {
return nil
}
func getExecutableFilePath() (string, error) {
+6
View File
@@ -23,3 +23,9 @@ func Println(enabled bool, v ...interface{}) {
Logger.Infof("NEZHA@%s>> %v", time.Now().Format("2006-01-02 15:04:05"), fmt.Sprint(v...))
}
}
func Printf(enabled bool, format string, v ...interface{}) {
if enabled {
Logger.Infof("NEZHA@%s>> "+format, append([]interface{}{time.Now().Format("2006-01-02 15:04:05")}, v...)...)
}
}