feat: add file transfer support (#55)
* feat: add file transfer support * 1MB buffer
This commit is contained in:
+82
-25
@@ -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++ {
|
||||
|
||||
Reference in New Issue
Block a user