feat: 添加支持文件路径处理函数,增强证书目录路径验证和日志记录功能

This commit is contained in:
ryan
2026-03-13 16:01:52 +08:00
parent 17f88917f4
commit 8c3dd75802
6 changed files with 290 additions and 12 deletions
+1 -1
View File
@@ -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
+30 -2
View File
@@ -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
+71 -6
View File
@@ -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)
}
}
+20
View File
@@ -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
}
+44 -3
View File
@@ -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
}
+124
View File
@@ -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")
}
}