feat(M4): 生命周期 + 审计 — 到期锁定/回收 cron、审计查询/CSV 导出/每日归档、settings 动态策略

This commit is contained in:
2026-08-30 10:17:22 +08:00
parent 1ea18490e0
commit c0e2ff975a
18 changed files with 1088 additions and 75 deletions
+299
View File
@@ -0,0 +1,299 @@
package service
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"time"
"ws_usernode/internal/model"
)
func TestSettingServiceCRUD(t *testing.T) {
db := testDB(t)
cfg := testConfig()
svc := NewSettingService(db, cfg)
ctx := context.Background()
// 初始值 = config 默认
items, err := svc.GetAll(ctx)
if err != nil {
t.Fatalf("getall: %v", err)
}
if len(items) != 3 {
t.Fatalf("items = %d, want 3", len(items))
}
for _, it := range items {
if it.Overridden {
t.Fatalf("initial item %s should not be overridden", it.Key)
}
}
if d := svc.DefaultTTL(ctx); d != cfg.Policy.DefaultTTL {
t.Fatalf("default ttl = %s, want %s", d, cfg.Policy.DefaultTTL)
}
// 更新
if err := svc.Set(ctx, "policy.default_ttl", "720h"); err != nil {
t.Fatalf("set: %v", err)
}
if d := svc.DefaultTTL(ctx); d != 720*time.Hour {
t.Fatalf("default ttl after set = %s, want 720h", d)
}
items, _ = svc.GetAll(ctx)
for _, it := range items {
if it.Key == "policy.default_ttl" {
if !it.Overridden || it.Value != "720h" {
t.Fatalf("item = %+v, want overridden 720h", it)
}
}
}
// 未知 key / 非法时长
if err := svc.Set(ctx, "smtp.host", "x"); !errors.Is(err, ErrSettingKeyUnknown) {
t.Fatalf("unknown key err = %v, want ErrSettingKeyUnknown", err)
}
if err := svc.Set(ctx, "policy.default_ttl", "not-a-duration"); !errors.Is(err, ErrSettingValueInvalid) {
t.Fatalf("bad value err = %v, want ErrSettingValueInvalid", err)
}
}
func TestSettingServiceRecycleRetention(t *testing.T) {
db := testDB(t)
cfg := testConfig()
svc := NewSettingService(db, cfg)
ctx := context.Background()
// 未覆盖时回退 config
if d := svc.RecyclePeriod(ctx); d != cfg.Policy.RecyclePeriod {
t.Fatalf("recycle = %s", d)
}
if d := svc.AuditRetention(ctx); d != cfg.Policy.AuditRetention {
t.Fatalf("retention = %s", d)
}
_ = svc.Set(ctx, "policy.recycle_period", "168h")
_ = svc.Set(ctx, "policy.audit_retention", "720h")
if d := svc.RecyclePeriod(ctx); d != 168*time.Hour {
t.Fatalf("recycle after set = %s", d)
}
}
func TestUserServiceDefaultTTLFromSettings(t *testing.T) {
db := testDB(t)
sys := newFakeSys()
cfg := testConfig()
users := NewUserService(db, sys, cfg)
settings := NewSettingService(db, cfg)
users.WithSettings(settings)
ctx := context.Background()
_ = settings.Set(ctx, "policy.default_ttl", "168h")
u, err := users.Create(ctx, "ttluser", "ttl@example.com", "", "", 0, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
if u.ExpireAt == nil {
t.Fatal("expire_at nil")
}
if d := time.Until(*u.ExpireAt); d < 160*time.Hour || d > 176*time.Hour {
t.Fatalf("expire in = %s, want ~168h", d)
}
}
func TestAuditServiceQueryExport(t *testing.T) {
db := testDB(t)
svc := NewAuditService(db, testLogger())
ctx := context.Background()
for i := 0; i < 5; i++ {
if err := svc.Record(ctx, 1, "root", "user.create", "user", "10", map[string]any{"n": i}, "127.0.0.1", model.ResultSuccess); err != nil {
t.Fatalf("record: %v", err)
}
}
if err := svc.Record(ctx, 2, "ext_x", "key.create", "ssh_key", "3", nil, "10.0.0.1", model.ResultSuccess); err != nil {
t.Fatalf("record: %v", err)
}
// 分页查询
rows, total, err := svc.Query(ctx, AuditFilter{Action: "user.create", Page: 1, PageSize: 2})
if err != nil {
t.Fatalf("query: %v", err)
}
if total != 5 || len(rows) != 2 {
t.Fatalf("query total=%d len=%d, want 5/2", total, len(rows))
}
// 按 actor 筛选
rows, total, _ = svc.Query(ctx, AuditFilter{ActorName: "ext_", Page: 1, PageSize: 10})
if total != 1 {
t.Fatalf("actor filter total = %d, want 1", total)
}
// CSV 导出(含 BOM 与表头)
var buf strings.Builder
if err := svc.ExportCSV(ctx, &buf, nil, nil); err != nil {
t.Fatalf("export: %v", err)
}
out := buf.String()
if !strings.HasPrefix(out, "\ufeffid,created_at,") {
t.Fatalf("csv missing header/BOM: %q", out[:40])
}
if !strings.Contains(out, "user.create") || !strings.Contains(out, "key.create") {
t.Fatalf("csv missing rows: %q", out)
}
}
func TestAuditServiceArchive(t *testing.T) {
db := testDB(t)
svc := NewAuditService(db, testLogger())
ctx := context.Background()
dir := t.TempDir()
// 3 条旧记录 + 1 条新记录
old := time.Now().Add(-10 * 24 * time.Hour)
for i := 0; i < 3; i++ {
if err := db.Create(&model.AuditLog{CreatedAt: old, Action: "old", ActorName: "root", Result: model.ResultSuccess}).Error; err != nil {
t.Fatalf("seed old: %v", err)
}
}
if err := svc.Record(ctx, 1, "root", "fresh", "user", "1", nil, "ip", model.ResultSuccess); err != nil {
t.Fatalf("seed fresh: %v", err)
}
n, err := svc.Archive(ctx, dir, time.Now().Add(-5*24*time.Hour))
if err != nil {
t.Fatalf("archive: %v", err)
}
if n != 3 {
t.Fatalf("archived = %d, want 3", n)
}
// 归档后旧记录已清理,新记录保留
var count int64
if err := db.Model(&model.AuditLog{}).Count(&count).Error; err != nil {
t.Fatalf("count: %v", err)
}
if count != 1 {
t.Fatalf("remaining = %d, want 1", count)
}
// 归档文件存在且含旧记录
files, _ := filepath.Glob(filepath.Join(dir, "audit-*.csv"))
if len(files) != 1 {
t.Fatalf("archive files = %v", files)
}
b, _ := os.ReadFile(files[0])
if !strings.Contains(string(b), "old") {
t.Fatalf("archive file missing rows: %q", string(b))
}
// dir 为空:跳过且不清理
n, err = svc.Archive(ctx, "", time.Now().Add(-5*24*time.Hour))
if err != nil || n != 0 {
t.Fatalf("empty dir archive = %d, %v, want 0/nil", n, err)
}
count = 0
if err := db.Model(&model.AuditLog{}).Count(&count).Error; err != nil {
t.Fatalf("count: %v", err)
}
if count != 1 {
t.Fatalf("remaining after empty-dir = %d, want 1(未配置目录不清理)", count)
}
}
func TestLifecycleScanExpired(t *testing.T) {
db := testDB(t)
sys := newFakeSys()
cfg := testConfig()
users := NewUserService(db, sys, cfg)
settings := NewSettingService(db, cfg)
mailer := &recordingMail{}
log := testLogger()
svc := NewLifecycleService(db, cfg, sys, users, settings, mailer, NewAuditService(db, log), log)
ctx := context.Background()
// 一个将过期、一个正常
u1, err := users.Create(ctx, "expire1", "e1@example.com", "", "", 90*24*time.Hour, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
u2, err := users.Create(ctx, "expire2", "e2@example.com", "", "", 90*24*time.Hour, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
past := time.Now().Add(-time.Hour)
if err := db.Model(&model.User{}).Where("id = ?", u1.ID).Update("expire_at", &past).Error; err != nil {
t.Fatalf("force expire u1: %v", err)
}
n, err := svc.ScanExpired(ctx)
if err != nil {
t.Fatalf("scan: %v", err)
}
if n != 1 {
t.Fatalf("expired = %d, want 1", n)
}
if got, _ := users.GetByID(ctx, u1.ID); got.Status != model.UserStatusExpired {
t.Fatalf("u1 status = %s, want expired", got.Status)
}
if got, _ := users.GetByID(ctx, u2.ID); got.Status != model.UserStatusActive {
t.Fatalf("u2 status = %s, want active", got.Status)
}
if mailer.lastTo != "e1@example.com" || !strings.Contains(mailer.lastBody, "到期") {
t.Fatalf("expire mail = %s / %s", mailer.lastTo, mailer.lastBody)
}
// 幂等:再次扫描不重复处理(已 expired)
if n, _ := svc.ScanExpired(ctx); n != 0 {
t.Fatalf("rescan = %d, want 0", n)
}
}
func TestLifecycleRecycle(t *testing.T) {
db := testDB(t)
sys := newFakeSys()
cfg := testConfig()
users := NewUserService(db, sys, cfg)
settings := NewSettingService(db, cfg)
mailer := &recordingMail{}
log := testLogger()
svc := NewLifecycleService(db, cfg, sys, users, settings, mailer, NewAuditService(db, log), log)
ctx := context.Background()
// 已过期很久(超回收期)与刚过期(回收期内)
u1, err := users.Create(ctx, "oldone", "o1@example.com", "", "", 90*24*time.Hour, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
u2, err := users.Create(ctx, "freshone", "f1@example.com", "", "", 90*24*time.Hour, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
longPast := time.Now().Add(-100 * 24 * time.Hour)
justPast := time.Now().Add(-time.Hour)
if err := db.Model(&model.User{}).Where("id = ?", u1.ID).Updates(map[string]any{"status": model.UserStatusExpired, "expire_at": &longPast}).Error; err != nil {
t.Fatalf("force u1: %v", err)
}
if err := db.Model(&model.User{}).Where("id = ?", u2.ID).Updates(map[string]any{"status": model.UserStatusExpired, "expire_at": &justPast}).Error; err != nil {
t.Fatalf("force u2: %v", err)
}
n, err := svc.Recycle(ctx)
if err != nil {
t.Fatalf("recycle: %v", err)
}
if n != 1 {
t.Fatalf("recycled = %d, want 1", n)
}
// 超期账号已删除(DB + 系统账号),回收期内账号保留
if _, err := users.GetByID(ctx, u1.ID); !errors.Is(err, ErrUserNotFound) {
t.Fatalf("u1 err = %v, want ErrUserNotFound", err)
}
if sys.has("ext_oldone") {
t.Fatal("u1 system account should be removed")
}
if _, err := users.GetByID(ctx, u2.ID); err != nil {
t.Fatalf("u2 should remain: %v", err)
}
if mailer.lastTo != "o1@example.com" || !strings.Contains(mailer.lastBody, "回收") {
t.Fatalf("recycle mail = %s / %s", mailer.lastTo, mailer.lastBody)
}
}