Files
rill/internal/config/config_test.go
T
kevin f440bb22b6 增加工具模块与可信代理配置
- internal/utils:ClientIP 按序枚举 CDN/代理头(含 RFC 7239 Forwarded),仅在可信代理来源时采信,否则回退直连 IP;RemoteIP 取直连地址;RandomString 生成安全随机串
- 新增 server.trusted_proxies 配置(IP/CIDR,ConfigVersion 2→3 自动补全),启动时同步应用到 gin 与 utils
- 初始管理员密码生成改用 utils.RandomString,原密码测试迁至 utils
2026-09-21 20:28:38 +08:00

211 lines
5.5 KiB
Go

package config
import (
"bytes"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"github.com/goccy/go-yaml"
)
func TestUpgradeConfigFillsMissing(t *testing.T) {
input := `server:
host: "127.0.0.1" # 自定义监听地址
port: 9000
custom:
keep: true
`
result, err := upgradeConfig([]byte(input))
if err != nil {
t.Fatalf("upgradeConfig 失败: %v", err)
}
if result == nil {
t.Fatal("期望产生补全结果")
}
if result.Version != 0 {
t.Fatalf("补全前版本 = %d, 期望 0", result.Version)
}
out := string(result.Data)
for _, want := range []string{
fmt.Sprintf("version: %d", ConfigVersion),
`host: "127.0.0.1"`,
"# 自定义监听地址",
"sock:",
`"web.sock"`,
"# unix socket 文件路径",
"mode: release",
"custom:",
"keep: true",
"api:",
`prefix: "/api"`,
"auth:",
`token_ttl: "24h"`,
"trusted_proxies",
} {
if !strings.Contains(out, want) {
t.Errorf("补全结果缺少 %q\n---\n%s", want, out)
}
}
var cfg Config
if err := yaml.Unmarshal(result.Data, &cfg); err != nil {
t.Fatalf("补全结果无法解析: %v\n%s", err, out)
}
if cfg.Version != ConfigVersion {
t.Errorf("version = %d, 期望 %d", cfg.Version, ConfigVersion)
}
if cfg.Server.Host != "127.0.0.1" || cfg.Server.Port != 9000 {
t.Errorf("用户已有值被覆盖: %+v", cfg.Server)
}
if cfg.Server.Sock != "web.sock" || cfg.Server.Mode != "release" {
t.Errorf("缺失项未按默认值补全: %+v", cfg.Server)
}
if len(result.Added) == 0 {
t.Error("期望记录新增配置项路径")
}
}
func TestUpgradeConfigIdempotent(t *testing.T) {
first, err := upgradeConfig([]byte("server:\n port: 9000\n"))
if err != nil || first == nil {
t.Fatalf("首次补全失败: result=%v err=%v", first, err)
}
second, err := upgradeConfig(first.Data)
if err != nil {
t.Fatalf("二次检查失败: %v", err)
}
if second != nil {
t.Errorf("版本已是最新,不应再次变更:\n%s", second.Data)
}
}
func TestUpgradeConfigSkipsCurrentVersion(t *testing.T) {
input := fmt.Sprintf("version: %d\nserver:\n host: \"0.0.0.0\"\n", ConfigVersion)
result, err := upgradeConfig([]byte(input))
if err != nil {
t.Fatalf("upgradeConfig 失败: %v", err)
}
if result != nil {
t.Errorf("版本一致时不应扫描补全:\n%s", result.Data)
}
}
func TestUpgradeConfigNewerVersion(t *testing.T) {
input := "version: 99\nserver:\n port: 9000\n"
result, err := upgradeConfig([]byte(input))
if err != nil {
t.Fatalf("upgradeConfig 失败: %v", err)
}
if result != nil {
t.Errorf("高版本配置不应被改写:\n%s", result.Data)
}
}
func TestUpgradeConfigNullSection(t *testing.T) {
input := "version: 0\napi:\nserver:\n host: \"0.0.0.0\"\n"
result, err := upgradeConfig([]byte(input))
if err != nil || result == nil {
t.Fatalf("补全失败: result=%v err=%v", result, err)
}
out := string(result.Data)
for _, want := range []string{`prefix: "/api"`, `max_age: "12h"`, "allow_origins"} {
if !strings.Contains(out, want) {
t.Errorf("空节未按默认子树补全,缺少 %q\n---\n%s", want, out)
}
}
var cfg Config
if err := yaml.Unmarshal(result.Data, &cfg); err != nil {
t.Fatalf("补全结果无法解析: %v", err)
}
if cfg.API.CORS.MaxAge != "12h" {
t.Errorf("api.cors.max_age = %q, 期望 12h", cfg.API.CORS.MaxAge)
}
}
func TestLoadConfigUpgradeWritesFile(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "config.yaml")
original := "server:\n host: \"127.0.0.1\"\n port: 9000\n"
if err := os.WriteFile(path, []byte(original), 0o600); err != nil {
t.Fatal(err)
}
cfg, err := LoadConfig(path)
if err != nil {
t.Fatalf("LoadConfig 失败: %v", err)
}
if cfg.Server.Port != 9000 || cfg.Server.Sock != "web.sock" {
t.Errorf("加载结果不符合预期: %+v", cfg.Server)
}
updated, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(updated), fmt.Sprintf("version: %d", ConfigVersion)) {
t.Errorf("磁盘配置未补全 version:\n%s", updated)
}
bak, err := os.ReadFile(path + ".bak")
if err != nil {
t.Fatalf("未生成备份文件: %v", err)
}
if string(bak) != original {
t.Errorf("备份内容与升级前不一致:\n%s", bak)
}
info, err := os.Stat(path)
if err != nil {
t.Fatal(err)
}
if info.Mode().Perm() != 0o600 {
t.Errorf("文件权限 = %o, 期望 600", info.Mode().Perm())
}
if _, err := LoadConfig(path); err != nil {
t.Fatalf("二次加载失败: %v", err)
}
after, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(updated, after) {
t.Errorf("版本已是最新,二次加载不应改写文件:\n--- before\n%s\n--- after\n%s", updated, after)
}
}
func TestLoadConfigGeneratesDefault(t *testing.T) {
path := filepath.Join(t.TempDir(), "data", "config.yaml")
cfg, err := LoadConfig(path)
if err != nil {
t.Fatalf("LoadConfig 失败: %v", err)
}
if cfg.Version != ConfigVersion || cfg.Server.Port != 8080 || cfg.Server.Sock != "web.sock" {
t.Errorf("默认配置不符合预期: version=%d server=%+v", cfg.Version, cfg.Server)
}
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("默认配置文件未生成: %v", err)
}
if !strings.Contains(string(data), fmt.Sprintf("version: %d", ConfigVersion)) {
t.Errorf("默认配置文件缺少 version:\n%s", data)
}
}
func TestDefaultTemplateVersionMatchesConst(t *testing.T) {
m, err := defaultMapping()
if err != nil {
t.Fatalf("解析默认模板失败: %v", err)
}
if version := versionOf(m); version != ConfigVersion {
t.Fatalf("默认模板 version = %d, ConfigVersion = %d,请同步更新", version, ConfigVersion)
}
}