diff --git a/.gitignore b/.gitignore index 60598b2a..e5bbc8ef 100644 --- a/.gitignore +++ b/.gitignore @@ -7,7 +7,7 @@ upload build *.log logs - +.gocache # If you prefer the allow list template instead of the deny list, see community template: # https://github.com/github/gitignore/blob/main/community/Golang/Go.AllowList.gitignore diff --git a/atsf_agent/internal/nginx/manager.go b/atsf_agent/internal/nginx/manager.go index 0b757f6e..d0bb00da 100644 --- a/atsf_agent/internal/nginx/manager.go +++ b/atsf_agent/internal/nginx/manager.go @@ -499,7 +499,10 @@ func (m *Manager) restore(state *backupState) error { return err } for _, file := range state.Files { - targetPath := filepath.Join(m.CertDir, filepath.Clean(file.Path)) + targetPath, err := m.supportFileTargetPath(file.Path) + if err != nil { + return err + } if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { return err } @@ -521,7 +524,10 @@ func (m *Manager) writeSupportFiles(supportFiles []protocol.SupportFile) error { return err } for _, file := range supportFiles { - targetPath := filepath.Join(m.CertDir, filepath.Clean(file.Path)) + targetPath, err := m.supportFileTargetPath(file.Path) + if err != nil { + return err + } if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { return err } @@ -573,6 +579,28 @@ func (m *Manager) readSupportFiles() ([]protocol.SupportFile, error) { return files, nil } +func (m *Manager) supportFileTargetPath(relativePath string) (string, error) { + if strings.TrimSpace(m.CertDir) == "" { + return "", errors.New("cert dir 不能为空") + } + normalizedPath := filepath.Clean(filepath.FromSlash(strings.TrimSpace(relativePath))) + if normalizedPath == "." || normalizedPath == "" { + return "", errors.New("support file path 不能为空") + } + if filepath.IsAbs(normalizedPath) || filepath.VolumeName(normalizedPath) != "" { + return "", fmt.Errorf("support file path %q must be relative", relativePath) + } + targetPath := filepath.Join(m.CertDir, normalizedPath) + relativeToBase, err := filepath.Rel(m.CertDir, targetPath) + if err != nil { + return "", err + } + if relativeToBase == ".." || strings.HasPrefix(relativeToBase, ".."+string(os.PathSeparator)) { + return "", fmt.Errorf("support file path %q escapes cert dir", relativePath) + } + return targetPath, nil +} + func (m *Manager) renderRouteConfig(content string) string { if m.NginxCertDir == "" { return content diff --git a/atsf_agent/internal/nginx/manager_test.go b/atsf_agent/internal/nginx/manager_test.go index 7411d664..c1fb25f4 100644 --- a/atsf_agent/internal/nginx/manager_test.go +++ b/atsf_agent/internal/nginx/manager_test.go @@ -6,6 +6,7 @@ import ( "os" "path/filepath" "reflect" + "runtime" "strings" "testing" @@ -204,14 +205,17 @@ func TestDockerExecutorStartsStoppedContainer(t *testing.T) { } func TestDockerExecutorRunContainerMountsManagedFiles(t *testing.T) { + mainConfigPath := filepath.Clean("/tmp/managed/nginx.conf") + routeConfigDir := filepath.Clean("/tmp/managed/conf.d") + certDir := filepath.Clean("/tmp/managed/certs") runner := &fakeRunner{} executor := &DockerExecutor{ DockerBinary: "docker", ContainerName: "atsflare-openresty", Image: "openresty/openresty:alpine", - MainConfigPath: filepath.Clean("/tmp/managed/nginx.conf"), - RouteConfigDir: filepath.Clean("/tmp/managed/conf.d"), - CertDir: filepath.Clean("/tmp/managed/certs"), + MainConfigPath: mainConfigPath, + RouteConfigDir: routeConfigDir, + CertDir: certDir, NginxCertDir: "/etc/nginx/atsflare-certs", Runner: runner, } @@ -229,9 +233,9 @@ func TestDockerExecutorRunContainerMountsManagedFiles(t *testing.T) { "--name", "atsflare-openresty", "-p", "80:80", "-p", "443:443", - "-v", "/tmp/managed/nginx.conf:" + DockerMainConfigPath, - "-v", "/tmp/managed/conf.d:/etc/nginx/conf.d", - "-v", "/tmp/managed/certs:/etc/nginx/atsflare-certs", + "-v", mainConfigPath + ":" + DockerMainConfigPath, + "-v", routeConfigDir + ":/etc/nginx/conf.d", + "-v", certDir + ":/etc/nginx/atsflare-certs", "openresty/openresty:alpine", } if !reflect.DeepEqual(runner.calls[0].args, expectedArgs) { @@ -537,3 +541,64 @@ func TestManagerRollbackRestoresSupportFiles(t *testing.T) { t.Fatalf("expected cert rollback, got %s", string(certData)) } } + +func TestManagerSupportFileTargetPathRejectsEscapes(t *testing.T) { + manager := &Manager{CertDir: filepath.Join(t.TempDir(), "certs")} + if err := os.MkdirAll(manager.CertDir, 0o755); err != nil { + t.Fatalf("MkdirAll failed: %v", err) + } + + absolutePath := "/tmp/evil.crt" + if runtime.GOOS == "windows" { + absolutePath = `C:/tmp/evil.crt` + } + + testCases := []struct { + path string + shouldErr bool + }{ + {path: "nested/1.crt", shouldErr: false}, + {path: "../escape.crt", shouldErr: true}, + {path: "..\\escape.crt", shouldErr: true}, + {path: absolutePath, shouldErr: true}, + {path: "", shouldErr: true}, + } + + for _, testCase := range testCases { + targetPath, err := manager.supportFileTargetPath(testCase.path) + if testCase.shouldErr { + if err == nil { + t.Fatalf("expected path %q to be rejected, got target %q", testCase.path, targetPath) + } + continue + } + if err != nil { + t.Fatalf("expected path %q to be accepted: %v", testCase.path, err) + } + if !strings.HasPrefix(targetPath, manager.CertDir) { + t.Fatalf("expected target path %q to stay under %q", targetPath, manager.CertDir) + } + } +} + +func TestManagerApplyRejectsSupportFilePathTraversal(t *testing.T) { + tempDir := t.TempDir() + manager := &Manager{ + MainConfigPath: filepath.Join(tempDir, "nginx.conf"), + RouteConfigPath: filepath.Join(tempDir, "routes.conf"), + CertDir: filepath.Join(tempDir, "certs"), + NginxCertDir: "/etc/nginx/atsflare-certs", + Executor: &fakeExecutor{}, + } + + err := manager.Apply(context.Background(), "main", "route", []protocol.SupportFile{ + {Path: "../escape.crt", Content: "bad"}, + }) + if err == nil { + t.Fatal("expected Apply to reject traversal path") + } + + if _, statErr := os.Stat(filepath.Join(tempDir, "escape.crt")); !os.IsNotExist(statErr) { + t.Fatalf("expected escaped file to not exist, stat err = %v", statErr) + } +} diff --git a/atsf_server/model/apply_log.go b/atsf_server/model/apply_log.go index 5a660ab9..6fb6af09 100644 --- a/atsf_server/model/apply_log.go +++ b/atsf_server/model/apply_log.go @@ -29,3 +29,23 @@ func GetLatestApplyLog(nodeID string) (*ApplyLog, error) { err := DB.Where("node_id = ?", nodeID).Order("id desc").First(log).Error return log, err } + +func GetLatestApplyLogsByNodeIDs(nodeIDs []string) (map[string]*ApplyLog, error) { + result := make(map[string]*ApplyLog) + if len(nodeIDs) == 0 { + return result, nil + } + + var logs []*ApplyLog + subQuery := DB.Model(&ApplyLog{}). + Select("MAX(id) AS id"). + Where("node_id IN ?", nodeIDs). + Group("node_id") + if err := DB.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil { + return nil, err + } + for _, log := range logs { + result[log.NodeID] = log + } + return result, nil +} diff --git a/atsf_server/service/agent.go b/atsf_server/service/agent.go index e39b30a2..4c301a4e 100644 --- a/atsf_server/service/agent.go +++ b/atsf_server/service/agent.go @@ -118,6 +118,7 @@ func HeartbeatNode(node *model.Node, payload AgentNodePayload) (*HeartbeatRespon if err := validateAgentNodePayload(payload); err != nil { return nil, err } + previous := *node updateNow := node.UpdateRequested restartOpenrestyNow := node.RestartOpenrestyRequested updateChannel := normalizeReleaseChannel(node.UpdateChannel) @@ -127,8 +128,11 @@ func HeartbeatNode(node *model.Node, payload AgentNodePayload) (*HeartbeatRespon node.UpdateChannel = ReleaseChannelStable.String() node.UpdateTag = "" node.RestartOpenrestyRequested = false - if err := model.DB.Model(node).Select("ip", "agent_version", "nginx_version", "openresty_status", "openresty_message", "status", "current_version", "last_seen_at", "last_error", "update_requested", "update_channel", "update_tag", "restart_openresty_requested").Updates(node).Error; err != nil { - return nil, err + changes := collectNodeHeartbeatChanges(&previous, node) + if len(changes) > 0 { + if err := model.DB.Model(node).Updates(changes).Error; err != nil { + return nil, err + } } activeConfig, err := GetActiveConfigMetaForAgent() if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { @@ -252,6 +256,14 @@ func ListNodeViews() ([]*NodeView, error) { if err != nil { return nil, err } + nodeIDs := make([]string, 0, len(nodes)) + for _, node := range nodes { + nodeIDs = append(nodeIDs, node.NodeID) + } + latestLogs, err := model.GetLatestApplyLogsByNodeIDs(nodeIDs) + if err != nil { + return nil, err + } views := make([]*NodeView, 0, len(nodes)) for _, node := range nodes { computedStatus := computeNodeStatus(node) @@ -266,7 +278,7 @@ func ListNodeViews() ([]*NodeView, error) { } view := buildNodeView(node) view.Status = computedStatus - if log, err := model.GetLatestApplyLog(node.NodeID); err == nil { + if log, ok := latestLogs[node.NodeID]; ok { view.LatestApplyResult = log.Result view.LatestApplyMessage = log.Message view.LatestApplyChecksum = log.Checksum @@ -300,3 +312,32 @@ func computeNodeStatus(node *model.Node) string { } return NodeStatusOnline } + +func collectNodeHeartbeatChanges(previous *model.Node, current *model.Node) map[string]any { + if previous == nil || current == nil { + return map[string]any{} + } + changes := make(map[string]any) + appendIfChanged := func(key string, before any, after any) { + if before != after { + changes[key] = after + } + } + appendIfChanged("name", previous.Name, current.Name) + appendIfChanged("ip", previous.IP, current.IP) + appendIfChanged("agent_version", previous.AgentVersion, current.AgentVersion) + appendIfChanged("nginx_version", previous.NginxVersion, current.NginxVersion) + appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus) + appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage) + appendIfChanged("status", previous.Status, current.Status) + appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion) + appendIfChanged("last_error", previous.LastError, current.LastError) + appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested) + appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel) + appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag) + appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested) + if !previous.LastSeenAt.Equal(current.LastSeenAt) { + changes["last_seen_at"] = current.LastSeenAt + } + return changes +} diff --git a/atsf_server/service/node_update_test.go b/atsf_server/service/node_update_test.go index f6a6d914..a804f26f 100644 --- a/atsf_server/service/node_update_test.go +++ b/atsf_server/service/node_update_test.go @@ -5,8 +5,10 @@ import ( "atsflare/model" "io" "net/http" + "sort" "strings" "testing" + "time" ) type roundTripFunc func(req *http.Request) (*http.Response, error) @@ -164,3 +166,125 @@ func TestRequestNodeOpenrestyRestart(t *testing.T) { t.Fatal("expected restart_openresty_requested to be true") } } + +func TestListNodeViewsIncludesLatestApplyLogsForMultipleNodes(t *testing.T) { + setupServiceTestDB(t) + + now := time.Now() + nodes := []*model.Node{ + { + NodeID: "node-a", + Name: "edge-a", + IP: "10.0.0.11", + AgentToken: "token-a", + AgentVersion: "v0.5.0", + NginxVersion: "1.27.1.2", + Status: NodeStatusOnline, + LastSeenAt: now, + }, + { + NodeID: "node-b", + Name: "edge-b", + IP: "10.0.0.12", + AgentToken: "token-b", + AgentVersion: "v0.5.0", + NginxVersion: "1.27.1.2", + Status: NodeStatusOnline, + LastSeenAt: now, + }, + } + for _, node := range nodes { + if err := node.Insert(); err != nil { + t.Fatalf("failed to insert node %s: %v", node.NodeID, err) + } + } + + logs := []*model.ApplyLog{ + {NodeID: "node-a", Version: "20260313-001", Result: ApplyResultOK, Message: "first success", CreatedAt: now.Add(-2 * time.Minute)}, + {NodeID: "node-a", Version: "20260313-002", Result: ApplyResultFailed, Message: "latest failure", CreatedAt: now.Add(-1 * time.Minute)}, + {NodeID: "node-b", Version: "20260313-003", Result: ApplyResultOK, Message: "latest success", CreatedAt: now}, + } + for _, log := range logs { + if err := model.DB.Create(log).Error; err != nil { + t.Fatalf("failed to insert apply log for %s: %v", log.NodeID, err) + } + } + + views, err := ListNodeViews() + if err != nil { + t.Fatalf("ListNodeViews failed: %v", err) + } + if len(views) != 2 { + t.Fatalf("expected 2 node views, got %d", len(views)) + } + + sort.Slice(views, func(i int, j int) bool { + return views[i].NodeID < views[j].NodeID + }) + + if views[0].NodeID != "node-a" || views[0].LatestApplyResult != ApplyResultFailed || views[0].LatestApplyMessage != "latest failure" { + t.Fatalf("unexpected latest apply log for node-a: %+v", views[0]) + } + if views[1].NodeID != "node-b" || views[1].LatestApplyResult != ApplyResultOK || views[1].LatestApplyMessage != "latest success" { + t.Fatalf("unexpected latest apply log for node-b: %+v", views[1]) + } +} + +func TestCollectNodeHeartbeatChangesOnlyReturnsChangedFields(t *testing.T) { + now := time.Now() + before := &model.Node{ + Name: "edge-1", + IP: "10.0.0.8", + AgentVersion: "v0.5.0", + NginxVersion: "1.27.1.2", + OpenrestyStatus: OpenrestyStatusHealthy, + OpenrestyMessage: "", + Status: NodeStatusOnline, + CurrentVersion: "20260313-001", + LastSeenAt: now.Add(-time.Minute), + LastError: "", + UpdateRequested: true, + UpdateChannel: "preview", + UpdateTag: "v0.5.0-rc.1", + RestartOpenrestyRequested: true, + } + after := &model.Node{ + Name: "edge-1", + IP: "10.0.0.8", + AgentVersion: "v0.5.0", + NginxVersion: "1.27.1.2", + OpenrestyStatus: OpenrestyStatusHealthy, + OpenrestyMessage: "", + Status: NodeStatusOnline, + CurrentVersion: "20260313-001", + LastSeenAt: now, + LastError: "", + UpdateRequested: false, + UpdateChannel: "stable", + UpdateTag: "", + RestartOpenrestyRequested: false, + } + + changes := collectNodeHeartbeatChanges(before, after) + if len(changes) != 5 { + t.Fatalf("expected 5 changed fields, got %d: %#v", len(changes), changes) + } + if _, ok := changes["last_seen_at"]; !ok { + t.Fatal("expected last_seen_at change to be included") + } + if value, ok := changes["update_requested"]; !ok || value != false { + t.Fatalf("expected update_requested reset, got %#v", value) + } + if value, ok := changes["update_channel"]; !ok || value != "stable" { + t.Fatalf("expected update_channel reset, got %#v", value) + } + if value, ok := changes["update_tag"]; !ok || value != "" { + t.Fatalf("expected update_tag reset, got %#v", value) + } + if value, ok := changes["restart_openresty_requested"]; !ok || value != false { + t.Fatalf("expected restart_openresty_requested reset, got %#v", value) + } + if _, ok := changes["ip"]; ok { + t.Fatal("did not expect unchanged ip to be included") + } +}