diff --git a/openflare_agent/internal/nginx/manager.go b/openflare_agent/internal/nginx/manager.go index b5d1d01a..6c5ecb51 100644 --- a/openflare_agent/internal/nginx/manager.go +++ b/openflare_agent/internal/nginx/manager.go @@ -367,25 +367,23 @@ func (m *Manager) EnsureLuaAssets() error { if strings.TrimSpace(m.LuaDir) == "" { return nil } - if err := os.RemoveAll(m.LuaDir); err != nil && !os.IsNotExist(err) { - return err - } - if err := os.MkdirAll(m.LuaDir, 0o755); err != nil { - return err - } + files := make([]managedFile, 0, len(ManagedObservabilityLuaFiles())) for _, file := range ManagedObservabilityLuaFiles() { targetPath, err := luaFileTargetPath(m.LuaDir, file.Path) if err != nil { return err } - if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { - return err - } - if err := os.WriteFile(targetPath, []byte(file.Content), 0o644); err != nil { + relativePath, err := filepath.Rel(m.LuaDir, targetPath) + if err != nil { return err } + files = append(files, managedFile{ + Path: filepath.ToSlash(relativePath), + Content: []byte(file.Content), + Mode: 0o644, + }) } - return nil + return syncManagedFiles(m.LuaDir, files) } func (m *Manager) EnsureRuntime(ctx context.Context, recreate bool) error { @@ -614,6 +612,12 @@ type backupState struct { Files []protocol.SupportFile } +type managedFile struct { + Path string + Content []byte + Mode fs.FileMode +} + func (m *Manager) backup() (*backupState, error) { if m.MainConfigPath == "" { return nil, errors.New("main config path 不能为空") @@ -678,50 +682,34 @@ func (m *Manager) restore(state *backupState) error { if m.CertDir == "" { return nil } - if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { - return err - } - if err := os.MkdirAll(m.CertDir, 0o755); err != nil { - return err - } - for _, file := range state.Files { - targetPath, err := m.certFileTargetPath(file.Path) - if err != nil { - return err - } - if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { - return err - } - if err := os.WriteFile(targetPath, []byte(file.Content), certFileMode(file.Path)); err != nil { - return err - } - } - return nil + return m.writeManagedCertFiles(state.Files) } func (m *Manager) writeCertFiles(certFiles []protocol.SupportFile) error { if m.CertDir == "" { return nil } - if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { - return err - } - if err := os.MkdirAll(m.CertDir, 0o755); err != nil { - return err - } + return m.writeManagedCertFiles(certFiles) +} + +func (m *Manager) writeManagedCertFiles(certFiles []protocol.SupportFile) error { + files := make([]managedFile, 0, len(certFiles)) for _, file := range certFiles { targetPath, err := m.certFileTargetPath(file.Path) if err != nil { return err } - if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { - return err - } - if err := os.WriteFile(targetPath, []byte(file.Content), certFileMode(file.Path)); err != nil { + relativePath, err := filepath.Rel(m.CertDir, targetPath) + if err != nil { return err } + files = append(files, managedFile{ + Path: filepath.ToSlash(relativePath), + Content: []byte(file.Content), + Mode: certFileMode(file.Path), + }) } - return nil + return syncManagedFiles(m.CertDir, files) } func (m *Manager) readCertFiles() ([]protocol.SupportFile, error) { @@ -828,6 +816,94 @@ func luaFileTargetPath(baseDir string, relativePath string) (string, error) { return targetPath, nil } +func syncManagedFiles(baseDir string, files []managedFile) error { + if strings.TrimSpace(baseDir) == "" { + return errors.New("managed dir cannot be empty") + } + if info, err := os.Stat(baseDir); err == nil && !info.IsDir() { + return fmt.Errorf("managed dir %q is not a directory", baseDir) + } else if err != nil && !os.IsNotExist(err) { + return err + } + if err := os.MkdirAll(baseDir, 0o755); err != nil { + return err + } + + desired := make(map[string]managedFile, len(files)) + for _, file := range files { + cleanPath := filepath.Clean(filepath.FromSlash(strings.TrimSpace(file.Path))) + if cleanPath == "." || cleanPath == "" { + return errors.New("managed file path cannot be empty") + } + desired[cleanPath] = managedFile{ + Path: cleanPath, + Content: file.Content, + Mode: file.Mode, + } + } + + if err := filepath.Walk(baseDir, func(path string, info os.FileInfo, err error) error { + if err != nil { + return err + } + if info.IsDir() { + return nil + } + relativePath, err := filepath.Rel(baseDir, path) + if err != nil { + return err + } + if _, ok := desired[filepath.Clean(relativePath)]; ok { + return nil + } + return os.Remove(path) + }); err != nil { + return err + } + + for _, file := range desired { + targetPath := filepath.Join(baseDir, file.Path) + if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { + return err + } + if err := os.WriteFile(targetPath, file.Content, file.Mode); err != nil { + return err + } + } + + return removeEmptyManagedDirs(baseDir) +} + +func removeEmptyManagedDirs(baseDir string) error { + dirs := make([]string, 0) + if err := filepath.Walk(baseDir, func(path string, info os.FileInfo, err error) error { + if err != nil { + return err + } + if info.IsDir() && path != baseDir { + dirs = append(dirs, path) + } + return nil + }); err != nil { + return err + } + sort.Slice(dirs, func(i int, j int) bool { + return len(dirs[i]) > len(dirs[j]) + }) + for _, dir := range dirs { + entries, err := os.ReadDir(dir) + if err != nil { + return err + } + if len(entries) == 0 { + if err := os.Remove(dir); err != nil && !os.IsNotExist(err) { + return err + } + } + } + return nil +} + func (m *Manager) renderRouteConfig(content string) string { if m.NginxCertDir == "" { return content diff --git a/openflare_agent/internal/nginx/manager_test.go b/openflare_agent/internal/nginx/manager_test.go index 8cfe135d..96950a6e 100644 --- a/openflare_agent/internal/nginx/manager_test.go +++ b/openflare_agent/internal/nginx/manager_test.go @@ -679,6 +679,61 @@ func TestResolverDirectiveForDockerMode(t *testing.T) { } } +func TestWriteCertFilesKeepsBaseDirAndRemovesStaleFiles(t *testing.T) { + tempDir := t.TempDir() + certDir := filepath.Join(tempDir, "certs") + if err := os.MkdirAll(filepath.Join(certDir, "stale"), 0o755); err != nil { + t.Fatalf("MkdirAll failed: %v", err) + } + if err := os.WriteFile(filepath.Join(certDir, "stale", "old.crt"), []byte("old"), 0o644); err != nil { + t.Fatalf("WriteFile failed: %v", err) + } + manager := &Manager{CertDir: certDir} + + if err := manager.writeCertFiles([]protocol.SupportFile{ + {Path: "1.crt", Content: "cert"}, + {Path: "1.key", Content: "key"}, + }); err != nil { + t.Fatalf("writeCertFiles failed: %v", err) + } + + if _, err := os.Stat(certDir); err != nil { + t.Fatalf("expected cert dir to persist, stat err = %v", err) + } + if _, err := os.Stat(filepath.Join(certDir, "stale", "old.crt")); !os.IsNotExist(err) { + t.Fatalf("expected stale cert file to be removed, stat err = %v", err) + } + if _, err := os.Stat(filepath.Join(certDir, "1.crt")); err != nil { + t.Fatalf("expected new cert file to exist, stat err = %v", err) + } +} + +func TestEnsureLuaAssetsKeepsBaseDirAndRemovesStaleFiles(t *testing.T) { + tempDir := t.TempDir() + luaDir := filepath.Join(tempDir, "lua") + if err := os.MkdirAll(filepath.Join(luaDir, "stale"), 0o755); err != nil { + t.Fatalf("MkdirAll failed: %v", err) + } + if err := os.WriteFile(filepath.Join(luaDir, "stale", "old.lua"), []byte("old"), 0o644); err != nil { + t.Fatalf("WriteFile failed: %v", err) + } + manager := &Manager{LuaDir: luaDir} + + if err := manager.EnsureLuaAssets(); err != nil { + t.Fatalf("EnsureLuaAssets failed: %v", err) + } + + if _, err := os.Stat(luaDir); err != nil { + t.Fatalf("expected lua dir to persist, stat err = %v", err) + } + if _, err := os.Stat(filepath.Join(luaDir, "stale", "old.lua")); !os.IsNotExist(err) { + t.Fatalf("expected stale lua file to be removed, stat err = %v", err) + } + if _, err := os.Stat(filepath.Join(luaDir, "log.lua")); err != nil { + t.Fatalf("expected managed lua file to exist, stat err = %v", err) + } +} + func TestCertFileMode(t *testing.T) { testCases := []struct { path string