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
+82 -25
View File
@@ -31,6 +31,7 @@ import (
"google.golang.org/grpc/credentials/insecure"
"github.com/nezhahq/agent/model"
fm "github.com/nezhahq/agent/pkg/fm"
"github.com/nezhahq/agent/pkg/monitor"
"github.com/nezhahq/agent/pkg/processgroup"
"github.com/nezhahq/agent/pkg/pty"
@@ -161,7 +162,7 @@ func init() {
func main() {
if err := agentCmd.Execute(); err != nil {
fmt.Println(err)
println(err)
os.Exit(1)
}
}
@@ -215,7 +216,11 @@ func run() {
// 下载远程命令执行需要的终端
if !agentCliParam.DisableCommandExecute {
go pty.DownloadDependency()
go func() {
if err := pty.DownloadDependency(); err != nil {
printf("pty 下载依赖失败: %v", err)
}
}()
}
// 上报服务器信息
go reportStateDaemon()
@@ -259,7 +264,7 @@ func run() {
}
conn, err = grpc.DialContext(timeOutCtx, agentCliParam.Server, securityOption, grpc.WithPerRPCCredentials(&auth))
if err != nil {
println("与面板建立连接失败", err)
printf("与面板建立连接失败: %v", err)
cancel()
retry()
continue
@@ -270,7 +275,7 @@ func run() {
timeOutCtx, cancel = context.WithTimeout(context.Background(), networkTimeOut)
_, err = client.ReportSystemInfo(timeOutCtx, monitor.GetHost().PB())
if err != nil {
println("上报系统信息失败", err)
printf("上报系统信息失败: %v", err)
cancel()
retry()
continue
@@ -280,12 +285,12 @@ func run() {
// 执行 Task
tasks, err := client.RequestTask(context.Background(), monitor.GetHost().PB())
if err != nil {
println("请求任务失败", err)
printf("请求任务失败: %v", err)
retry()
continue
}
err = receiveTasks(tasks)
println("receiveTasks exit to main", err)
printf("receiveTasks exit to main: %v", err)
retry()
}
}
@@ -293,7 +298,7 @@ func run() {
func runService(action string, flags []string) {
dir, err := os.Getwd()
if err != nil {
println("获取当前工作目录时出错: ", err)
printf("获取当前工作目录时出错: ", err)
return
}
@@ -315,7 +320,7 @@ func runService(action string, flags []string) {
}
s, err := service.New(prg, svcConfig)
if err != nil {
log.Printf("创建服务时出错,以普通模式运行: %v", err)
printf("创建服务时出错,以普通模式运行: %v", err)
run()
return
}
@@ -324,7 +329,7 @@ func runService(action string, flags []string) {
if agentConfig.Debug {
serviceLogger, err := s.Logger(nil)
if err != nil {
log.Printf("获取 service logger 时出错: %+v", err)
printf("获取 service logger 时出错: %+v", err)
} else {
util.Logger = serviceLogger
}
@@ -332,7 +337,7 @@ func runService(action string, flags []string) {
if action == "install" {
initName := s.Platform()
log.Println("Init system is:", initName)
println("Init system is:", initName)
}
if len(action) != 0 {
@@ -351,7 +356,7 @@ func runService(action string, flags []string) {
func receiveTasks(tasks pb.NezhaService_RequestTaskClient) error {
var err error
defer println("receiveTasks exit", time.Now(), "=>", err)
defer printf("receiveTasks exit %v => %v", time.Now(), err)
for {
var task *pb.Task
task, err = tasks.Recv()
@@ -393,10 +398,13 @@ func doTask(task *pb.Task) {
case model.TaskTypeReportHostInfo:
reportState(time.Time{})
return
case model.TaskTypeFM:
handleFMTask(task)
return
case model.TaskTypeKeepalive:
return
default:
println("不支持的任务", task)
printf("不支持的任务: %v", task)
return
}
client.ReportTask(context.Background(), &result)
@@ -406,7 +414,7 @@ func doTask(task *pb.Task) {
func reportStateDaemon() {
var lastReportHostInfo time.Time
var err error
defer println("reportState exit", time.Now(), "=>", err)
defer printf("reportState exit %v => %v", time.Now(), err)
for {
// 为了更准确的记录时段流量,inited 后再上传状态信息
lastReportHostInfo = reportState(lastReportHostInfo)
@@ -421,7 +429,7 @@ func reportState(lastReportHostInfo time.Time) time.Time {
_, err := client.ReportSystemState(timeOutCtx, monitor.GetState(agentCliParam.SkipConnectionCount, agentCliParam.SkipProcsCount).PB())
cancel()
if err != nil {
println("reportState error", err)
printf("reportState error: %v", err)
time.Sleep(delayWhenError)
}
// 每10分钟重新获取一次硬件信息
@@ -445,7 +453,7 @@ func doSelfUpdate(useLocalVersion bool) {
if useLocalVersion {
v = semver.MustParse(version)
}
println("检查更新", v)
printf("检查更新: %v", v)
var latest *selfupdate.Release
var err error
if monitor.CachedCountryCode != "cn" && !agentCliParam.UseGiteeToUpgrade {
@@ -454,11 +462,11 @@ func doSelfUpdate(useLocalVersion bool) {
latest, err = selfupdate.UpdateSelfGitee(v, "naibahq/agent")
}
if err != nil {
println("更新失败", err)
printf("更新失败: %v", err)
return
}
if !latest.Version.Equals(v) {
println("已经更新至", latest.Version, " 正在结束进程")
printf("已经更新至: %v, 正在结束进程", latest.Version)
os.Exit(1)
}
}
@@ -668,13 +676,13 @@ func handleTerminalTask(task *pb.Task) {
var terminal model.TerminalTask
err := util.Json.Unmarshal([]byte(task.GetData()), &terminal)
if err != nil {
println("Terminal 任务解析错误", err)
printf("Terminal 任务解析错误: %v", err)
return
}
remoteIO, err := client.IOStream(context.Background())
if err != nil {
println("Terminal IOStream失败", err)
printf("Terminal IOStream失败: %v", err)
return
}
@@ -682,13 +690,13 @@ func handleTerminalTask(task *pb.Task) {
if err := remoteIO.Send(&pb.IOStreamData{Data: append([]byte{
0xff, 0x05, 0xff, 0x05,
}, []byte(terminal.StreamID)...)}); err != nil {
println("Terminal 发送StreamID失败", err)
printf("Terminal 发送StreamID失败: %v", err)
return
}
tty, err := pty.Start()
if err != nil {
println("Terminal pty.Start失败", err)
printf("Terminal pty.Start失败 %v", err)
return
}
@@ -739,13 +747,13 @@ func handleNATTask(task *pb.Task) {
var nat model.TaskNAT
err := util.Json.Unmarshal([]byte(task.GetData()), &nat)
if err != nil {
println("NAT 任务解析错误", err)
printf("NAT 任务解析错误: %v", err)
return
}
remoteIO, err := client.IOStream(context.Background())
if err != nil {
println("NAT IOStream失败", err)
printf("NAT IOStream失败: %v", err)
return
}
@@ -753,13 +761,13 @@ func handleNATTask(task *pb.Task) {
if err := remoteIO.Send(&pb.IOStreamData{Data: append([]byte{
0xff, 0x05, 0xff, 0x05,
}, []byte(nat.StreamID)...)}); err != nil {
println("NAT 发送StreamID失败", err)
printf("NAT 发送StreamID失败: %v", err)
return
}
conn, err := net.Dial("tcp", nat.Host)
if err != nil {
println(fmt.Sprintf("NAT Dial %s 失败:%s", nat.Host, err))
printf("NAT Dial %s 失败:%s", nat.Host, err)
return
}
@@ -792,10 +800,59 @@ func handleNATTask(task *pb.Task) {
}
}
func handleFMTask(task *pb.Task) {
if agentCliParam.DisableCommandExecute {
println("此 Agent 已禁止命令执行")
return
}
var fmTask model.TaskFM
err := util.Json.Unmarshal([]byte(task.GetData()), &fmTask)
if err != nil {
printf("FM 任务解析错误: %v", err)
return
}
remoteIO, err := client.IOStream(context.Background())
if err != nil {
printf("FM IOStream失败: %v", err)
return
}
// 发送 StreamID
if err := remoteIO.Send(&pb.IOStreamData{Data: append([]byte{
0xff, 0x05, 0xff, 0x05,
}, []byte(fmTask.StreamID)...)}); err != nil {
printf("FM 发送StreamID失败: %v", err)
return
}
defer func() {
errCloseSend := remoteIO.CloseSend()
println("FM exit", fmTask.StreamID, nil, errCloseSend)
}()
println("FM init", fmTask.StreamID)
fmc := fm.NewFMClient(remoteIO, printf)
for {
var remoteData *pb.IOStreamData
if remoteData, err = remoteIO.Recv(); err != nil {
return
}
if remoteData.Data == nil || len(remoteData.Data) == 0 {
return
}
fmc.DoTask(remoteData)
}
}
func println(v ...interface{}) {
util.Println(agentConfig.Debug, v...)
}
func printf(format string, v ...interface{}) {
util.Printf(agentConfig.Debug, format, v...)
}
func generateQueue(start int, size int) []int {
var result []int
for i := start; i < start+size; i++ {