mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-28 05:46:36 +08:00
feat: 添加支持文件路径处理函数,增强证书目录路径验证和日志记录功能
This commit is contained in:
+1
-1
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user