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