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 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) }