package config import ( "os" "path/filepath" "testing" "time" ) func writeTemp(t *testing.T, content string) string { t.Helper() dir := t.TempDir() path := filepath.Join(dir, "config.toml") if err := os.WriteFile(path, []byte(content), 0o600); err != nil { t.Fatal(err) } return path } func TestLoadDefault(t *testing.T) { cfg, err := LoadDefault() if err != nil { t.Fatalf("LoadDefault: %v", err) } if cfg.Database.Driver != "sqlite" { t.Errorf("default driver = %q, want sqlite", cfg.Database.Driver) } if cfg.Policy.DefaultTTL != 90*24*time.Hour { t.Errorf("default ttl = %v", cfg.Policy.DefaultTTL) } if cfg.System.UserPrefix != "ext_" { t.Errorf("default user_prefix = %q", cfg.System.UserPrefix) } } func TestLoadFileOverrides(t *testing.T) { path := writeTemp(t, ` [server] listen = "0.0.0.0:9999" session_ttl = "2h" [policy] default_ttl = "720h" `) cfg, err := Load(path) if err != nil { t.Fatalf("Load: %v", err) } if cfg.Server.Listen != "0.0.0.0:9999" { t.Errorf("listen = %q", cfg.Server.Listen) } if cfg.Server.SessionTTL != 2*time.Hour { t.Errorf("session_ttl = %v", cfg.Server.SessionTTL) } if cfg.Policy.DefaultTTL != 30*24*time.Hour { t.Errorf("default_ttl = %v", cfg.Policy.DefaultTTL) } } func TestEnvOverrides(t *testing.T) { t.Setenv("USERNODE_DATABASE_DRIVER", "mysql") t.Setenv("USERNODE_DATABASE_DSN", "u:p@tcp(h:3306)/db") t.Setenv("USERNODE_POLICY_OTPTTL", "5m") t.Setenv("USERNODE_SYSTEM_SUDO", "true") t.Setenv("USERNODE_SERVER_TRUSTED_PROXIES", "10.0.0.1, 10.0.0.2") cfg, err := LoadDefault() if err != nil { t.Fatalf("LoadDefault: %v", err) } if cfg.Database.Driver != "mysql" || cfg.Database.DSN != "u:p@tcp(h:3306)/db" { t.Errorf("database = %+v", cfg.Database) } if cfg.Policy.OTPTTL != 5*time.Minute { t.Errorf("otp_ttl = %v", cfg.Policy.OTPTTL) } if !cfg.System.Sudo { t.Error("system.sudo should be true") } if len(cfg.Server.TrustedProxies) != 2 || cfg.Server.TrustedProxies[0] != "10.0.0.1" { t.Errorf("trusted_proxies = %v", cfg.Server.TrustedProxies) } } func TestInvalidDriver(t *testing.T) { path := writeTemp(t, "[database]\ndriver = \"oracle\"\ndsn = \"x\"\n") if _, err := Load(path); err == nil { t.Fatal("expected error for unsupported driver") } } func TestCamelToSnake(t *testing.T) { cases := map[string]string{ "SessionTTL": "session_ttl", "Listen": "listen", "OTPTTL": "otpttl", // 全大写缩写按一个词处理(与 strcase 行为一致) "OTPCooldown": "otp_cooldown", "TrustedProxies": "trusted_proxies", "AuthorizedKeysDir": "authorized_keys_dir", "HomeBase": "home_base", } for in, want := range cases { if got := camelToSnake(in); got != want { t.Errorf("camelToSnake(%q) = %q, want %q", in, got, want) } } }