feat: enhance path management and improve backup logic for certificate directory

This commit is contained in:
ryan
2026-03-10 11:30:26 +08:00
parent 87ab1f8664
commit 2185326f2f
4 changed files with 51 additions and 9 deletions
+37 -5
View File
@@ -4,7 +4,9 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"os" "os"
pathpkg "path"
"path/filepath" "path/filepath"
"strings"
"time" "time"
) )
@@ -70,18 +72,18 @@ func applyDefaults(cfg *Config, baseDir string) {
cfg.DataDir = filepath.Join(baseDir, "data") cfg.DataDir = filepath.Join(baseDir, "data")
} }
if cfg.NginxPath == "" { if cfg.NginxPath == "" {
cfg.RouteConfigPath = filepath.Join(cfg.DataDir, defaultDockerRouteConfigRelativePath) cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath)
cfg.StatePath = filepath.Join(cfg.DataDir, defaultDockerStateRelativePath) cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath)
} else { } else {
if cfg.RouteConfigPath == "" { if cfg.RouteConfigPath == "" {
cfg.RouteConfigPath = filepath.Join(cfg.DataDir, defaultDockerRouteConfigRelativePath) cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath)
} }
if cfg.StatePath == "" { if cfg.StatePath == "" {
cfg.StatePath = filepath.Join(cfg.DataDir, defaultDockerStateRelativePath) cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath)
} }
} }
if cfg.CertDir == "" { if cfg.CertDir == "" {
cfg.CertDir = filepath.Join(cfg.DataDir, defaultCertDirRelativePath) cfg.CertDir = joinManagedPath(cfg.DataDir, defaultCertDirRelativePath)
} }
if cfg.NginxCertDir == "" { if cfg.NginxCertDir == "" {
if cfg.NginxPath != "" { if cfg.NginxPath != "" {
@@ -99,6 +101,36 @@ func applyDefaults(cfg *Config, baseDir string) {
if cfg.RequestTimeout <= 0 { if cfg.RequestTimeout <= 0 {
cfg.RequestTimeout = 10 * time.Second cfg.RequestTimeout = 10 * time.Second
} }
normalizeManagedPaths(cfg)
}
func normalizeManagedPaths(cfg *Config) {
if cfg == nil {
return
}
if usesSlashPath(cfg.DataDir) {
cfg.DataDir = filepath.ToSlash(cfg.DataDir)
}
if usesSlashPath(cfg.RouteConfigPath) {
cfg.RouteConfigPath = filepath.ToSlash(cfg.RouteConfigPath)
}
if usesSlashPath(cfg.CertDir) {
cfg.CertDir = filepath.ToSlash(cfg.CertDir)
}
if usesSlashPath(cfg.StatePath) {
cfg.StatePath = filepath.ToSlash(cfg.StatePath)
}
}
func usesSlashPath(path string) bool {
return strings.HasPrefix(path, "/")
}
func joinManagedPath(base string, relative string) string {
if usesSlashPath(base) {
return pathpkg.Join(filepath.ToSlash(base), relative)
}
return filepath.Join(base, relative)
} }
func validate(cfg *Config) error { func validate(cfg *Config) error {
+10 -2
View File
@@ -254,8 +254,10 @@ func (m *Manager) backup() (*backupState, error) {
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil { if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil {
return nil, err return nil, err
} }
if err := os.MkdirAll(m.CertDir, 0o755); err != nil { if m.CertDir != "" {
return nil, err if err := os.MkdirAll(m.CertDir, 0o755); err != nil {
return nil, err
}
} }
state := &backupState{} state := &backupState{}
data, err := os.ReadFile(m.RouteConfigPath) data, err := os.ReadFile(m.RouteConfigPath)
@@ -284,6 +286,9 @@ func (m *Manager) restore(state *backupState) error {
} else if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) { } else if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) {
return err return err
} }
if m.CertDir == "" {
return nil
}
if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) {
return err return err
} }
@@ -303,6 +308,9 @@ func (m *Manager) restore(state *backupState) error {
} }
func (m *Manager) writeSupportFiles(supportFiles []protocol.SupportFile) error { func (m *Manager) writeSupportFiles(supportFiles []protocol.SupportFile) error {
if m.CertDir == "" {
return nil
}
if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) {
return err return err
} }
+2 -1
View File
@@ -27,7 +27,8 @@ func printHelp() {
} }
func init() { func init() {
if !strings.HasSuffix(os.Args[0], ".test") { executableName := strings.ToLower(filepath.Base(os.Args[0]))
if !strings.Contains(executableName, ".test") {
flag.Parse() flag.Parse()
} }
+2 -1
View File
@@ -14,8 +14,8 @@ import (
"github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie" "github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"mime/multipart"
"math/big" "math/big"
"mime/multipart"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"path/filepath" "path/filepath"
@@ -213,6 +213,7 @@ func setupTestDB(t *testing.T) {
t.Helper() t.Helper()
dbPath := filepath.Join(t.TempDir(), "phase1.db") dbPath := filepath.Join(t.TempDir(), "phase1.db")
common.SQLitePath = dbPath common.SQLitePath = dbPath
common.AgentToken = "phase1-agent-token"
if err := model.InitDB(); err != nil { if err := model.InitDB(); err != nil {
t.Fatalf("failed to init db: %v", err) t.Fatalf("failed to init db: %v", err)
} }