Files
2026-07-22 16:51:23 +08:00

187 lines
5.9 KiB
Go

package service
import (
"context"
"errors"
"testing"
)
// patchOf 把完整设置转为全字段补丁,等价旧全量保存,供既有用例复用。
func patchOf(in SecuritySettings) SecurityPatch {
return SecurityPatch{
LoginFailLimit: &in.LoginFailLimit,
LoginLockMinutes: &in.LoginLockMinutes,
IPRateRPS: &in.IPRateRPS,
IPRateBurst: &in.IPRateBurst,
RealIPHeader: &in.RealIPHeader,
AppURL: &in.AppURL,
}
}
func TestSecurityReadWrite(t *testing.T) {
tests := []struct {
name string
setup func(t *testing.T, svc *SettingService)
want SecuritySettings
}{
{
name: "键全部缺省回退默认",
setup: func(t *testing.T, svc *SettingService) {},
want: securityDefaults,
},
{
name: "写入后读回",
setup: func(t *testing.T, svc *SettingService) {
in := SecuritySettings{LoginFailLimit: 10, LoginLockMinutes: 30, IPRateRPS: 20, IPRateBurst: 60, RealIPHeader: "X-Client-IP", AppURL: "https://demo.example.com/"}
if err := svc.UpdateSecurity(context.Background(), patchOf(in)); err != nil {
t.Fatalf("UpdateSecurity: %v", err)
}
},
want: SecuritySettings{LoginFailLimit: 10, LoginLockMinutes: 30, IPRateRPS: 20, IPRateBurst: 60, RealIPHeader: "X-Client-IP", AppURL: "https://demo.example.com"},
},
{
name: "脏数据按字段回退默认",
setup: func(t *testing.T, svc *SettingService) {
for k, v := range map[string]string{
settingSecLoginFailLimit: "abc",
settingSecIPRateRPS: "9999",
settingSecRealIPHeader: "X Evil Header\r\n",
} {
if err := svc.set(context.Background(), k, v); err != nil {
t.Fatalf("set: %v", err)
}
}
},
want: securityDefaults,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, _ := newSettingEnv(t)
tt.setup(t, svc)
got, err := svc.Security(context.Background())
if err != nil {
t.Fatalf("Security: %v", err)
}
if got != tt.want {
t.Errorf("settings = %+v, want %+v", got, tt.want)
}
})
}
}
func TestUpdateSecurityRejectsInvalid(t *testing.T) {
base := securityDefaults
tests := []struct {
name string
mutate func(*SecuritySettings)
}{
{name: "失败阈值越界", mutate: func(s *SecuritySettings) { s.LoginFailLimit = 0 }},
{name: "锁定时长越界", mutate: func(s *SecuritySettings) { s.LoginLockMinutes = 1441 }},
{name: "限速越界", mutate: func(s *SecuritySettings) { s.IPRateRPS = 101 }},
{name: "突发越界", mutate: func(s *SecuritySettings) { s.IPRateBurst = 0 }},
{name: "请求头含非法字符", mutate: func(s *SecuritySettings) { s.RealIPHeader = "X Custom;" }},
{name: "面板地址缺协议", mutate: func(s *SecuritySettings) { s.AppURL = "demo.example.com" }},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, _ := newSettingEnv(t)
in := base
tt.mutate(&in)
if err := svc.UpdateSecurity(context.Background(), patchOf(in)); !errors.Is(err, ErrInvalidSecurity) {
t.Fatalf("err = %v, want ErrInvalidSecurity", err)
}
})
}
}
func TestSecurityCachedSnapshot(t *testing.T) {
svc, _ := newSettingEnv(t)
// 未初始化时返回默认,不触库
if got := svc.SecurityCached(); got != securityDefaults {
t.Errorf("cold cache = %+v, want defaults", got)
}
if err := svc.UpdateSecurity(context.Background(), SecurityPatch{LoginFailLimit: intPtr(8)}); err != nil {
t.Fatalf("UpdateSecurity: %v", err)
}
// Update 即刷新快照
if got := svc.SecurityCached(); got.LoginFailLimit != 8 {
t.Errorf("cached limit = %d, want 8", got.LoginFailLimit)
}
// nil 容忍
if got := securityOf(nil); got != securityDefaults {
t.Errorf("securityOf(nil) = %+v, want defaults", got)
}
}
func TestEffectiveAppURL(t *testing.T) {
svc, _ := newSettingEnv(t)
svc.SetEnvPublicURL("https://env.example.com/")
if got := svc.EffectiveAppURL(); got != "https://env.example.com" {
t.Errorf("env fallback = %q", got)
}
if err := svc.UpdateSecurity(context.Background(), SecurityPatch{AppURL: strPtr("https://app.example.com")}); err != nil {
t.Fatalf("UpdateSecurity: %v", err)
}
if got := svc.EffectiveAppURL(); got != "https://app.example.com" {
t.Errorf("app_url 应优先, got %q", got)
}
}
// TestUpdateSecurityPartialPatch 锁定 PATCH 语义:只写出现字段,其余不回滚。
func TestUpdateSecurityPartialPatch(t *testing.T) {
seeded := SecuritySettings{LoginFailLimit: 10, LoginLockMinutes: 30, IPRateRPS: 20, IPRateBurst: 60, RealIPHeader: "X-Client-IP", AppURL: "https://a.example.com"}
tests := []struct {
name string
patch SecurityPatch
want SecuritySettings
wantErr bool
}{
{
name: "单字段更新不动其余",
patch: SecurityPatch{LoginFailLimit: intPtr(3)},
want: func() SecuritySettings {
w := seeded
w.LoginFailLimit = 3
return w
}(),
},
{name: "空补丁不改任何值", patch: SecurityPatch{}, want: seeded},
{
name: "补丁字段规范化(AppURL 去尾斜杠)",
patch: SecurityPatch{AppURL: strPtr("https://b.example.com/")},
want: func() SecuritySettings {
w := seeded
w.AppURL = "https://b.example.com"
return w
}(),
},
{name: "补丁字段越界拒绝", patch: SecurityPatch{IPRateRPS: intPtr(0)}, wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, _ := newSettingEnv(t)
if err := svc.UpdateSecurity(context.Background(), patchOf(seeded)); err != nil {
t.Fatalf("seed: %v", err)
}
err := svc.UpdateSecurity(context.Background(), tt.patch)
if tt.wantErr {
if !errors.Is(err, ErrInvalidSecurity) {
t.Fatalf("err = %v, want ErrInvalidSecurity", err)
}
return
}
if err != nil {
t.Fatalf("UpdateSecurity: %v", err)
}
got, err := svc.Security(context.Background())
if err != nil || got != tt.want {
t.Errorf("settings = %+v (%v), want %+v", got, err, tt.want)
}
if cached := svc.SecurityCached(); cached != tt.want {
t.Errorf("cached = %+v, want %+v", cached, tt.want)
}
})
}
}