diff --git a/atsf_agent/internal/config/config.go b/atsf_agent/internal/config/config.go index dd3e64dd..e192e6b1 100644 --- a/atsf_agent/internal/config/config.go +++ b/atsf_agent/internal/config/config.go @@ -4,7 +4,9 @@ import ( "encoding/json" "errors" "os" + pathpkg "path" "path/filepath" + "strings" "time" ) @@ -70,18 +72,18 @@ func applyDefaults(cfg *Config, baseDir string) { cfg.DataDir = filepath.Join(baseDir, "data") } if cfg.NginxPath == "" { - cfg.RouteConfigPath = filepath.Join(cfg.DataDir, defaultDockerRouteConfigRelativePath) - cfg.StatePath = filepath.Join(cfg.DataDir, defaultDockerStateRelativePath) + cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath) + cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath) } else { if cfg.RouteConfigPath == "" { - cfg.RouteConfigPath = filepath.Join(cfg.DataDir, defaultDockerRouteConfigRelativePath) + cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath) } if cfg.StatePath == "" { - cfg.StatePath = filepath.Join(cfg.DataDir, defaultDockerStateRelativePath) + cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath) } } if cfg.CertDir == "" { - cfg.CertDir = filepath.Join(cfg.DataDir, defaultCertDirRelativePath) + cfg.CertDir = joinManagedPath(cfg.DataDir, defaultCertDirRelativePath) } if cfg.NginxCertDir == "" { if cfg.NginxPath != "" { @@ -99,6 +101,36 @@ func applyDefaults(cfg *Config, baseDir string) { if cfg.RequestTimeout <= 0 { 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 { diff --git a/atsf_agent/internal/nginx/manager.go b/atsf_agent/internal/nginx/manager.go index 5673f9d8..11d02f6a 100644 --- a/atsf_agent/internal/nginx/manager.go +++ b/atsf_agent/internal/nginx/manager.go @@ -254,8 +254,10 @@ func (m *Manager) backup() (*backupState, error) { if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil { return nil, err } - if err := os.MkdirAll(m.CertDir, 0o755); err != nil { - return nil, err + if m.CertDir != "" { + if err := os.MkdirAll(m.CertDir, 0o755); err != nil { + return nil, err + } } state := &backupState{} 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) { return err } + if m.CertDir == "" { + return nil + } if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { return err } @@ -303,6 +308,9 @@ func (m *Manager) restore(state *backupState) 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) { return err } diff --git a/atsf_server/common/init.go b/atsf_server/common/init.go index c66090ca..7f926cc1 100644 --- a/atsf_server/common/init.go +++ b/atsf_server/common/init.go @@ -27,7 +27,8 @@ func printHelp() { } func init() { - if !strings.HasSuffix(os.Args[0], ".test") { + executableName := strings.ToLower(filepath.Base(os.Args[0])) + if !strings.Contains(executableName, ".test") { flag.Parse() } diff --git a/atsf_server/router/api_phase1_test.go b/atsf_server/router/api_phase1_test.go index b81cb61c..90e9cec2 100644 --- a/atsf_server/router/api_phase1_test.go +++ b/atsf_server/router/api_phase1_test.go @@ -14,8 +14,8 @@ import ( "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/cookie" "github.com/gin-gonic/gin" - "mime/multipart" "math/big" + "mime/multipart" "net/http" "net/http/httptest" "path/filepath" @@ -213,6 +213,7 @@ func setupTestDB(t *testing.T) { t.Helper() dbPath := filepath.Join(t.TempDir(), "phase1.db") common.SQLitePath = dbPath + common.AgentToken = "phase1-agent-token" if err := model.InitDB(); err != nil { t.Fatalf("failed to init db: %v", err) }