260 lines
7.1 KiB
Go
260 lines
7.1 KiB
Go
package config
|
||
|
||
import (
|
||
"os"
|
||
"path/filepath"
|
||
"strings"
|
||
"testing"
|
||
)
|
||
|
||
func backendRoot(t *testing.T) string {
|
||
t.Helper()
|
||
wd, err := os.Getwd()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
root, err := filepath.Abs(filepath.Join(wd, ".."))
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
return root
|
||
}
|
||
|
||
func assertPath(t *testing.T, got, want string) {
|
||
t.Helper()
|
||
g, err := filepath.Abs(got)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
w, err := filepath.Abs(want)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if filepath.Clean(g) != filepath.Clean(w) {
|
||
t.Fatalf("路径不符:\n got %s\n want %s", g, w)
|
||
}
|
||
}
|
||
|
||
func TestResolveWorkPathFromOfficialCwd(t *testing.T) {
|
||
root := backendRoot(t)
|
||
t.Chdir(root)
|
||
got, err := resolveWorkPath()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
assertPath(t, got, root)
|
||
}
|
||
|
||
func TestResolveWorkPathFromCmdJiang13(t *testing.T) {
|
||
root := backendRoot(t)
|
||
t.Chdir(filepath.Join(root, "cmd", "jiang13"))
|
||
got, err := resolveWorkPath()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
assertPath(t, got, root)
|
||
}
|
||
|
||
func TestResolveWorkPathFromRepoRoot(t *testing.T) {
|
||
root := backendRoot(t)
|
||
t.Chdir(filepath.Join(root, ".."))
|
||
got, err := resolveWorkPath()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
assertPath(t, got, root)
|
||
}
|
||
|
||
func TestResolveWorkPathEnvOverride(t *testing.T) {
|
||
root := backendRoot(t)
|
||
t.Setenv("JIANG13_WORK_PATH", root)
|
||
t.Chdir(t.TempDir())
|
||
got, err := resolveWorkPath()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
assertPath(t, got, root)
|
||
}
|
||
|
||
func TestRejectStrayDataDir(t *testing.T) {
|
||
err := rejectStrayDataDir(`C:\proj\backend\cmd\jiang13\data`)
|
||
if err == nil {
|
||
t.Fatal("期望拒绝 cmd/jiang13/data")
|
||
}
|
||
if !strings.Contains(err.Error(), "拒绝使用数据目录") {
|
||
t.Fatalf("错误文案不符: %v", err)
|
||
}
|
||
if err := rejectStrayDataDir(`C:\proj\backend\data`); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
}
|
||
|
||
func TestAllowOriginDevAndProd(t *testing.T) {
|
||
dev := &Config{DevMode: true}
|
||
if !dev.AllowOrigin("http://localhost:3000") || !dev.AllowOrigin("http://127.0.0.1:3000") {
|
||
t.Fatal("开发态应放行 next dev Origin")
|
||
}
|
||
if dev.AllowOrigin("https://bbs.example.com") {
|
||
t.Fatal("开发态不应放行生产域名")
|
||
}
|
||
|
||
prod := &Config{
|
||
DevMode: false,
|
||
SiteURL: "https://bbs.example.com",
|
||
CORSOrigins: []string{"https://mirror.example.com"},
|
||
}
|
||
if !prod.AllowOrigin("https://bbs.example.com/") {
|
||
t.Fatal("生产态应放行 SITE_URL")
|
||
}
|
||
if !prod.AllowOrigin("https://mirror.example.com") {
|
||
t.Fatal("生产态应放行 CORS_ORIGINS")
|
||
}
|
||
if prod.AllowOrigin("http://localhost:3000") {
|
||
t.Fatal("生产态不应放行 localhost")
|
||
}
|
||
}
|
||
|
||
func TestParseSettingsMasterKeyFromIni(t *testing.T) {
|
||
work := t.TempDir()
|
||
key := "AQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQE="
|
||
if err := os.WriteFile(filepath.Join(work, "app.ini"), []byte("[security]\nJWT_SECRET = test-secret-not-for-prod\nSETTINGS_MASTER_KEY = "+key+"\n"), 0600); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
t.Setenv("JIANG13_WORK_PATH", work)
|
||
t.Setenv("SETTINGS_MASTER_KEY", "")
|
||
cfg, err := Parse()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if cfg.SettingsMasterKey != key {
|
||
t.Fatalf("未从 app.ini 读取主密钥: %q", cfg.SettingsMasterKey)
|
||
}
|
||
|
||
override := "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA="
|
||
t.Setenv("SETTINGS_MASTER_KEY", override)
|
||
cfg, err = Parse()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if cfg.SettingsMasterKey != override {
|
||
t.Fatalf("环境变量应覆盖 app.ini: %q", cfg.SettingsMasterKey)
|
||
}
|
||
}
|
||
|
||
func TestParseSiteURLAndDataDir(t *testing.T) {
|
||
work := t.TempDir()
|
||
data := t.TempDir()
|
||
t.Setenv("JIANG13_WORK_PATH", work)
|
||
t.Setenv("DEV_MODE", "false")
|
||
t.Setenv("SITE_URL", "https://bbs.example.com/")
|
||
t.Setenv("CORS_ORIGINS", " https://a.example.com ,https://b.example.com/ ")
|
||
t.Setenv("DATA_DIR", data)
|
||
t.Setenv("JWT_SECRET", "test-secret-not-for-prod")
|
||
|
||
cfg, err := Parse()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if cfg.DevMode {
|
||
t.Fatal("DEV_MODE=false 应关闭开发态")
|
||
}
|
||
if cfg.SiteURL != "https://bbs.example.com" {
|
||
t.Fatalf("SITE_URL 未去尾斜杠: %q", cfg.SiteURL)
|
||
}
|
||
if len(cfg.CORSOrigins) != 2 || cfg.CORSOrigins[0] != "https://a.example.com" || cfg.CORSOrigins[1] != "https://b.example.com" {
|
||
t.Fatalf("CORS_ORIGINS 解析错误: %#v", cfg.CORSOrigins)
|
||
}
|
||
assertPath(t, cfg.DataDir, data)
|
||
}
|
||
|
||
func TestEnsureAppIniBackfillsMissingKeys(t *testing.T) {
|
||
work := t.TempDir()
|
||
iniPath := filepath.Join(work, "app.ini")
|
||
// 旧版精简配置:缺 SITE_URL / CORS_ORIGINS / SETTINGS_MASTER_KEY
|
||
old := "[server]\nHTTP_PORT = 3001\n\n[database]\nDSN = postgres://u:p@localhost/db\n\n[security]\nJWT_SECRET = keep-me\n\n[paths]\nDATA = data\n\n[app]\nDEV_MODE = true\n"
|
||
if err := os.WriteFile(iniPath, []byte(old), 0600); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
|
||
f, err := loadAndEnsureAppIni(iniPath)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !f.Section("app").HasKey("SITE_URL") || !f.Section("app").HasKey("CORS_ORIGINS") {
|
||
t.Fatal("应补全 SITE_URL / CORS_ORIGINS")
|
||
}
|
||
if got := f.Section("app").Key("SITE_URL").String(); got != "http://localhost:3000" {
|
||
t.Fatalf("开发态缺 SITE_URL 应补本机 origin,got %q", got)
|
||
}
|
||
if got := f.Section("security").Key("JWT_SECRET").String(); got != "keep-me" {
|
||
t.Fatalf("已有值被改写: %q", got)
|
||
}
|
||
if !f.Section("security").HasKey("SETTINGS_MASTER_KEY") {
|
||
t.Fatal("应补全 SETTINGS_MASTER_KEY 空键")
|
||
}
|
||
|
||
// 再生产态缺项:只留空键,不写 localhost
|
||
prod := t.TempDir()
|
||
prodIni := filepath.Join(prod, "app.ini")
|
||
if err := os.WriteFile(prodIni, []byte("[app]\nDEV_MODE = false\n"), 0600); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
pf, err := loadAndEnsureAppIni(prodIni)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if got := pf.Section("app").Key("SITE_URL").String(); got != "" {
|
||
t.Fatalf("生产态缺 SITE_URL 应留空提示填写,got %q", got)
|
||
}
|
||
}
|
||
|
||
func TestParseSiteURLFromIni(t *testing.T) {
|
||
work := t.TempDir()
|
||
body := "[security]\nJWT_SECRET = test-secret-not-for-prod\n\n[app]\nDEV_MODE = true\nSITE_URL = https://forum.example.com/\nCORS_ORIGINS = https://a.example.com, https://b.example.com/\n"
|
||
if err := os.WriteFile(filepath.Join(work, "app.ini"), []byte(body), 0600); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
t.Setenv("JIANG13_WORK_PATH", work)
|
||
t.Setenv("SITE_URL", "")
|
||
t.Setenv("CORS_ORIGINS", "")
|
||
t.Setenv("DEV_MODE", "")
|
||
t.Setenv("JWT_SECRET", "")
|
||
|
||
cfg, err := Parse()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if cfg.SiteURL != "https://forum.example.com" {
|
||
t.Fatalf("未从 app.ini 读取 SITE_URL: %q", cfg.SiteURL)
|
||
}
|
||
if len(cfg.CORSOrigins) != 2 || cfg.CORSOrigins[0] != "https://a.example.com" {
|
||
t.Fatalf("未从 app.ini 读取 CORS_ORIGINS: %#v", cfg.CORSOrigins)
|
||
}
|
||
|
||
t.Setenv("SITE_URL", "https://env.example.com")
|
||
cfg, err = Parse()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if cfg.SiteURL != "https://env.example.com" {
|
||
t.Fatalf("环境变量应覆盖 app.ini SITE_URL: %q", cfg.SiteURL)
|
||
}
|
||
}
|
||
|
||
func TestEnsureAppIniCreatesWhenMissing(t *testing.T) {
|
||
work := t.TempDir()
|
||
iniPath := filepath.Join(work, "app.ini")
|
||
f, err := loadAndEnsureAppIni(iniPath)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, err := os.Stat(iniPath); err != nil {
|
||
t.Fatal("应创建 app.ini")
|
||
}
|
||
for _, spec := range appIniSchema {
|
||
if !f.Section(spec.section).HasKey(spec.key) {
|
||
t.Fatalf("新建文件缺少 %s.%s", spec.section, spec.key)
|
||
}
|
||
}
|
||
}
|