feat(M4): 生命周期 + 审计 — 到期锁定/回收 cron、审计查询/CSV 导出/每日归档、settings 动态策略
This commit is contained in:
+153
-10
@@ -2,9 +2,13 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/csv"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -13,14 +17,15 @@ import (
|
||||
)
|
||||
|
||||
// AuditService 审计服务:append-only 记录(业务代码仅允许 INSERT),
|
||||
// 保留策略与 CSV 导出归档在 M4 完成,这里给出接口与最小实现。
|
||||
// 查询/CSV 导出/每日归档(PLAN F6)。
|
||||
type AuditService struct {
|
||||
db *gorm.DB
|
||||
db *gorm.DB
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
// NewAuditService 创建审计服务。
|
||||
func NewAuditService(db *gorm.DB) *AuditService {
|
||||
return &AuditService{db: db}
|
||||
func NewAuditService(db *gorm.DB, log *slog.Logger) *AuditService {
|
||||
return &AuditService{db: db, log: log}
|
||||
}
|
||||
|
||||
// Record 记录一条管理操作审计。detail 为任意结构体,入库前 JSON 序列化。
|
||||
@@ -42,10 +47,148 @@ func (s *AuditService) Record(ctx context.Context, actorID uint, actorName, acti
|
||||
return s.db.WithContext(ctx).Create(&entry).Error
|
||||
}
|
||||
|
||||
// ExportCSV 导出审计为 CSV。M0 骨架:返回占位错误,M4 实现手动导出 + 每日归档。
|
||||
// AuditFilter 审计查询条件。
|
||||
type AuditFilter struct {
|
||||
ActorName string // 操作者模糊匹配
|
||||
Action string // 动作精确匹配
|
||||
ResourceType string // 资源类型(user / ssh_key / approval ...)
|
||||
ResourceID string // 资源 ID
|
||||
Since *time.Time // 起始时间(含)
|
||||
Until *time.Time // 结束时间(含)
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// Query 分页查询审计日志(最新在前)。
|
||||
func (s *AuditService) Query(ctx context.Context, f AuditFilter) ([]model.AuditLog, int64, error) {
|
||||
q := s.db.WithContext(ctx).Model(&model.AuditLog{})
|
||||
if f.ActorName != "" {
|
||||
q = q.Where("actor_name LIKE ?", "%"+f.ActorName+"%")
|
||||
}
|
||||
if f.Action != "" {
|
||||
q = q.Where("action = ?", f.Action)
|
||||
}
|
||||
if f.ResourceType != "" {
|
||||
q = q.Where("resource_type = ?", f.ResourceType)
|
||||
}
|
||||
if f.ResourceID != "" {
|
||||
q = q.Where("resource_id = ?", f.ResourceID)
|
||||
}
|
||||
if f.Since != nil {
|
||||
q = q.Where("created_at >= ?", *f.Since)
|
||||
}
|
||||
if f.Until != nil {
|
||||
q = q.Where("created_at <= ?", *f.Until)
|
||||
}
|
||||
var total int64
|
||||
if err := q.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
page, size := f.Page, f.PageSize
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if size < 1 {
|
||||
size = 20
|
||||
}
|
||||
if size > 200 {
|
||||
size = 200
|
||||
}
|
||||
var rows []model.AuditLog
|
||||
if err := q.Order("id DESC").Offset((page - 1) * size).Limit(size).Find(&rows).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return rows, total, nil
|
||||
}
|
||||
|
||||
// ExportCSV 导出审计为 CSV(手动导出接口)。since/until 为空则导出全量。
|
||||
// 输出带 UTF-8 BOM,Excel 直接打开不乱码。
|
||||
func (s *AuditService) ExportCSV(ctx context.Context, w io.Writer, since, until *time.Time) error {
|
||||
_ = w
|
||||
_ = since
|
||||
_ = until
|
||||
return errors.New("service: 审计 CSV 导出将在 M4 实现")
|
||||
q := s.db.WithContext(ctx).Model(&model.AuditLog{})
|
||||
if since != nil {
|
||||
q = q.Where("created_at >= ?", *since)
|
||||
}
|
||||
if until != nil {
|
||||
q = q.Where("created_at <= ?", *until)
|
||||
}
|
||||
var rows []model.AuditLog
|
||||
if err := q.Order("id ASC").Find(&rows).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return writeAuditCSV(w, rows)
|
||||
}
|
||||
|
||||
// Archive 每日归档:将 before 之前的审计记录导出到 dir 后删除(PLAN F6
|
||||
// "归档后可安全清理")。dir 为空时跳过并返回 0(防止未配置目录就丢审计)。
|
||||
// 写文件失败时不删除任何记录(归档成功是清理的前提)。
|
||||
func (s *AuditService) Archive(ctx context.Context, dir string, before time.Time) (int, error) {
|
||||
if dir == "" {
|
||||
s.log.Warn("audit: archive_dir 未配置,跳过归档与清理(防止丢审计)")
|
||||
return 0, nil
|
||||
}
|
||||
var rows []model.AuditLog
|
||||
if err := s.db.WithContext(ctx).Where("created_at < ?", before).Order("id ASC").Find(&rows).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
if err := os.MkdirAll(dir, 0o750); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
// 按归档执行日期分文件,同日追加
|
||||
name := "audit-" + time.Now().Format("2006-01-02") + ".csv"
|
||||
f, err := os.OpenFile(filepath.Join(dir, name), os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o640)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := writeAuditCSV(f, rows); err != nil {
|
||||
f.Close()
|
||||
return 0, err
|
||||
}
|
||||
if err := f.Close(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
// 归档成功后清理(append-only 约束:清理是运维动作,走 cron 每日任务)
|
||||
if err := s.db.WithContext(ctx).Where("created_at < ?", before).Delete(&model.AuditLog{}).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
s.log.Info("audit: archived", "file", name, "count", len(rows), "before", before.Format(time.RFC3339))
|
||||
return len(rows), nil
|
||||
}
|
||||
|
||||
// writeAuditCSV 以固定列序写审计记录(ExportCSV 与 Archive 共用)。
|
||||
func writeAuditCSV(w io.Writer, rows []model.AuditLog) error {
|
||||
// UTF-8 BOM:Excel 识别中文
|
||||
if _, err := w.Write([]byte{0xEF, 0xBB, 0xBF}); err != nil {
|
||||
return err
|
||||
}
|
||||
cw := csv.NewWriter(w)
|
||||
header := []string{"id", "created_at", "actor_id", "actor_name", "action", "resource_type", "resource_id", "detail", "ip", "result"}
|
||||
if err := cw.Write(header); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, r := range rows {
|
||||
rec := []string{
|
||||
strconv.FormatUint(uint64(r.ID), 10),
|
||||
r.CreatedAt.Format(time.RFC3339),
|
||||
strconv.FormatUint(uint64(r.ActorID), 10),
|
||||
r.ActorName,
|
||||
r.Action,
|
||||
r.ResourceType,
|
||||
r.ResourceID,
|
||||
r.Detail,
|
||||
r.IP,
|
||||
r.Result,
|
||||
}
|
||||
if err := cw.Write(rec); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
cw.Flush()
|
||||
if err := cw.Error(); err != nil {
|
||||
return err
|
||||
}
|
||||
// 追加模式归档时,BOM 会在文件中部重复,无害(解析器按首列内容处理)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -37,7 +37,7 @@ func newTestAuthService(t *testing.T) (*AuthService, *fakeSys, *recordingMailer)
|
||||
sys := newFakeSys()
|
||||
userSvc := NewUserService(db, sys, testConfig())
|
||||
adminSvc := NewAdminService(db)
|
||||
auditSvc := NewAuditService(db)
|
||||
auditSvc := NewAuditService(db, testLogger())
|
||||
cfg := testConfig()
|
||||
|
||||
mailer := &recordingMailer{}
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"ws_usernode/internal/config"
|
||||
"ws_usernode/internal/mail"
|
||||
"ws_usernode/internal/model"
|
||||
"ws_usernode/internal/system"
|
||||
)
|
||||
|
||||
// LifecycleService 账号生命周期维护(PLAN F2):
|
||||
// - 到期:锁定系统账号 + 清空 authorized_keys(SSH 立即失效)+ 状态置 expired + 邮件提醒;
|
||||
// - 回收:超过回收期仍未延期 → 删除系统账号与密钥(保留审计)+ 邮件通知。
|
||||
//
|
||||
// 由 cron 每日触发(单实例部署,PLAN §11)。
|
||||
type LifecycleService struct {
|
||||
db *gorm.DB
|
||||
cfg *config.Config
|
||||
sys system.Manager
|
||||
users *UserService
|
||||
settings *SettingService
|
||||
mailer mail.Mailer
|
||||
audit *AuditService
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
// NewLifecycleService 创建生命周期服务。
|
||||
func NewLifecycleService(db *gorm.DB, cfg *config.Config, sys system.Manager, users *UserService, settings *SettingService, mailer mail.Mailer, audit *AuditService, log *slog.Logger) *LifecycleService {
|
||||
return &LifecycleService{db: db, cfg: cfg, sys: sys, users: users, settings: settings, mailer: mailer, audit: audit, log: log}
|
||||
}
|
||||
|
||||
// ScanExpired 扫描到期账号(expire_at < now 且状态非 disabled/expired):
|
||||
// 锁定系统账号 + 清空 authorized_keys + 状态置 expired + 邮件通知(到期提醒)。
|
||||
// 返回本次处理的账号数。
|
||||
func (s *LifecycleService) ScanExpired(ctx context.Context) (int, error) {
|
||||
now := time.Now()
|
||||
var rows []model.User
|
||||
if err := s.db.WithContext(ctx).
|
||||
Where("status = ? AND expire_at IS NOT NULL AND expire_at < ?", model.UserStatusActive, now).
|
||||
Find(&rows).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
processed := 0
|
||||
for i := range rows {
|
||||
u := &rows[i]
|
||||
if !s.systemAccountOK(u.Username) {
|
||||
s.log.Warn("lifecycle: 系统账号缺失,跳过锁定", "username", u.Username)
|
||||
continue
|
||||
}
|
||||
if err := s.sys.SetLock(ctx, u.Username, true); err != nil {
|
||||
s.log.Error("lifecycle: 锁定系统账号失败", "username", u.Username, "err", err)
|
||||
_ = s.audit.Record(ctx, 0, "system", "lifecycle.expire", "user", fmt.Sprint(u.ID), map[string]any{"err": err.Error()}, "cron", model.ResultFailed)
|
||||
continue
|
||||
}
|
||||
if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, nil); err != nil {
|
||||
s.log.Error("lifecycle: 清空 authorized_keys 失败", "username", u.Username, "err", err)
|
||||
_ = s.audit.Record(ctx, 0, "system", "lifecycle.expire", "user", fmt.Sprint(u.ID), map[string]any{"err": err.Error()}, "cron", model.ResultFailed)
|
||||
continue
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", u.ID).Update("status", model.UserStatusExpired).Error; err != nil {
|
||||
s.log.Error("lifecycle: 更新状态失败", "username", u.Username, "err", err)
|
||||
continue
|
||||
}
|
||||
processed++
|
||||
s.notifyExpired(ctx, u)
|
||||
_ = s.audit.Record(ctx, 0, "system", "lifecycle.expire", "user", fmt.Sprint(u.ID), nil, "cron", model.ResultSuccess)
|
||||
}
|
||||
return processed, nil
|
||||
}
|
||||
|
||||
// Recycle 回收超期账号:expired 状态且过期时间超过回收期 → 删除系统账号与密钥
|
||||
// (保留审计)+ 邮件通知(回收提醒)。返回本次回收的账号数。
|
||||
func (s *LifecycleService) Recycle(ctx context.Context) (int, error) {
|
||||
period := s.settings.RecyclePeriod(ctx)
|
||||
cutoff := time.Now().Add(-period)
|
||||
var rows []model.User
|
||||
if err := s.db.WithContext(ctx).
|
||||
Where("status = ? AND expire_at IS NOT NULL AND expire_at < ?", model.UserStatusExpired, cutoff).
|
||||
Find(&rows).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
recycled := 0
|
||||
for i := range rows {
|
||||
u := &rows[i]
|
||||
if err := s.users.Delete(ctx, u.ID); err != nil {
|
||||
s.log.Error("lifecycle: 回收账号失败", "username", u.Username, "err", err)
|
||||
_ = s.audit.Record(ctx, 0, "system", "lifecycle.recycle", "user", fmt.Sprint(u.ID), map[string]any{"err": err.Error()}, "cron", model.ResultFailed)
|
||||
continue
|
||||
}
|
||||
recycled++
|
||||
s.notifyRecycled(ctx, u)
|
||||
_ = s.audit.Record(ctx, 0, "system", "lifecycle.recycle", "user", fmt.Sprint(u.ID), nil, "cron", model.ResultSuccess)
|
||||
}
|
||||
return recycled, nil
|
||||
}
|
||||
|
||||
// systemAccountOK 与 UserService 一致:dry-run 模式跳过真实检查。
|
||||
func (s *LifecycleService) systemAccountOK(username string) bool {
|
||||
if s.cfg.System.DryRun {
|
||||
return true
|
||||
}
|
||||
ok, err := s.sys.Exists(context.Background(), username)
|
||||
return err == nil && ok
|
||||
}
|
||||
|
||||
// notifyExpired 到期提醒:账号已到期,回收期内可延期恢复。
|
||||
func (s *LifecycleService) notifyExpired(ctx context.Context, u *model.User) {
|
||||
body := fmt.Sprintf(`您的服务器账号(%s)已到期,系统账号已被锁定,SSH 登录已不可用。
|
||||
|
||||
如需继续使用,请在回收期(%s)内联系管理员延期恢复;超过回收期账号将被自动回收。
|
||||
`, u.Username, s.settings.RecyclePeriod(ctx).String())
|
||||
if err := s.mailer.Send(ctx, u.Email, "服务器账号已到期", body); err != nil {
|
||||
s.log.Warn("lifecycle: 到期邮件发送失败", "username", u.Username, "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
// notifyRecycled 回收提醒:账号已超回收期被回收。
|
||||
func (s *LifecycleService) notifyRecycled(ctx context.Context, u *model.User) {
|
||||
body := fmt.Sprintf(`您的服务器账号(%s)已超过回收期且未办理延期,账号已被回收删除(系统账号与密钥均已清理)。
|
||||
|
||||
如需重新使用,请重新提交账号申请。
|
||||
`, u.Username)
|
||||
if err := s.mailer.Send(ctx, u.Email, "服务器账号已回收", body); err != nil {
|
||||
s.log.Warn("lifecycle: 回收邮件发送失败", "username", u.Username, "err", err)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"ws_usernode/internal/config"
|
||||
"ws_usernode/internal/model"
|
||||
)
|
||||
|
||||
// 系统设置错误。
|
||||
var ErrSettingKeyUnknown = errors.New("service: 未知设置项")
|
||||
var ErrSettingValueInvalid = errors.New("service: 设置值不合法")
|
||||
|
||||
// 设置项 key 白名单(settings 表可覆盖 config 的策略字段,PLAN F7)。
|
||||
// 其余设置(SMTP、会话时长等)为启动时配置,不支持动态覆盖。
|
||||
var settingKeys = map[string]bool{
|
||||
"policy.default_ttl": true, // 新账号默认有效期(时长)
|
||||
"policy.recycle_period": true, // 到期后回收期(时长)
|
||||
"policy.audit_retention": true, // 审计保留时长(时长)
|
||||
}
|
||||
|
||||
// SettingItem 一个设置项的生效值。
|
||||
type SettingItem struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"` // 生效值(settings 覆盖优先,否则 config 默认)
|
||||
Overridden bool `json:"overridden"` // 是否被 settings 表覆盖
|
||||
}
|
||||
|
||||
// SettingService 系统设置读写(settings 表,config 提供默认值)。
|
||||
type SettingService struct {
|
||||
db *gorm.DB
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
// NewSettingService 创建设置服务。
|
||||
func NewSettingService(db *gorm.DB, cfg *config.Config) *SettingService {
|
||||
return &SettingService{db: db, cfg: cfg}
|
||||
}
|
||||
|
||||
// GetAll 返回全部设置项的生效值(未覆盖的展示 config 默认值)。
|
||||
func (s *SettingService) GetAll(ctx context.Context) ([]SettingItem, error) {
|
||||
var rows []model.Setting
|
||||
if err := s.db.WithContext(ctx).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
overridden := make(map[string]string, len(rows))
|
||||
for _, r := range rows {
|
||||
if settingKeys[r.Key] {
|
||||
overridden[r.Key] = r.Value
|
||||
}
|
||||
}
|
||||
items := make([]SettingItem, 0, len(settingKeys))
|
||||
for _, key := range sortedSettingKeys() {
|
||||
if v, ok := overridden[key]; ok {
|
||||
items = append(items, SettingItem{Key: key, Value: v, Overridden: true})
|
||||
} else {
|
||||
items = append(items, SettingItem{Key: key, Value: s.defaultValue(key), Overridden: false})
|
||||
}
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
// Set 更新一个设置项:校验 key 白名单与值格式(均为时长)。全部校验通过才写入。
|
||||
func (s *SettingService) Set(ctx context.Context, key, value string) error {
|
||||
if !settingKeys[key] {
|
||||
return fmt.Errorf("%w: %q", ErrSettingKeyUnknown, key)
|
||||
}
|
||||
value = strings.TrimSpace(value)
|
||||
if _, err := time.ParseDuration(value); err != nil {
|
||||
return fmt.Errorf("%w: %q(需为时长,如 2160h)", ErrSettingValueInvalid, value)
|
||||
}
|
||||
return s.db.WithContext(ctx).Save(&model.Setting{Key: key, Value: value}).Error
|
||||
}
|
||||
|
||||
// Duration 读取策略时长设置:settings 覆盖优先,无覆盖或解析失败用 fallback。
|
||||
func (s *SettingService) Duration(ctx context.Context, key string, fallback time.Duration) time.Duration {
|
||||
if !settingKeys[key] {
|
||||
return fallback
|
||||
}
|
||||
var row model.Setting
|
||||
err := s.db.WithContext(ctx).First(&row, "key = ?", key).Error
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
d, err := time.ParseDuration(row.Value)
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// DefaultTTL 新账号默认有效期(settings 覆盖 config.policy.default_ttl)。
|
||||
func (s *SettingService) DefaultTTL(ctx context.Context) time.Duration {
|
||||
return s.Duration(ctx, "policy.default_ttl", s.cfg.Policy.DefaultTTL)
|
||||
}
|
||||
|
||||
// RecyclePeriod 到期后回收期。
|
||||
func (s *SettingService) RecyclePeriod(ctx context.Context) time.Duration {
|
||||
return s.Duration(ctx, "policy.recycle_period", s.cfg.Policy.RecyclePeriod)
|
||||
}
|
||||
|
||||
// AuditRetention 审计保留时长。
|
||||
func (s *SettingService) AuditRetention(ctx context.Context) time.Duration {
|
||||
return s.Duration(ctx, "policy.audit_retention", s.cfg.Policy.AuditRetention)
|
||||
}
|
||||
|
||||
// defaultValue 返回设置项的 config 默认值(保证遍历顺序稳定)。
|
||||
func (s *SettingService) defaultValue(key string) string {
|
||||
switch key {
|
||||
case "policy.default_ttl":
|
||||
return s.cfg.Policy.DefaultTTL.String()
|
||||
case "policy.recycle_period":
|
||||
return s.cfg.Policy.RecyclePeriod.String()
|
||||
case "policy.audit_retention":
|
||||
return s.cfg.Policy.AuditRetention.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func sortedSettingKeys() []string {
|
||||
// 固定顺序展示,避免 map 遍历随机
|
||||
return []string{"policy.default_ttl", "policy.recycle_period", "policy.audit_retention"}
|
||||
}
|
||||
+31
-12
@@ -25,9 +25,10 @@ var (
|
||||
|
||||
// UserService 外部用户生命周期服务:DB 记录 + system.Manager 系统账号操作。
|
||||
type UserService struct {
|
||||
db *gorm.DB
|
||||
sys system.Manager
|
||||
cfg *config.Config
|
||||
db *gorm.DB
|
||||
sys system.Manager
|
||||
cfg *config.Config
|
||||
settings *SettingService // 可选:默认 TTL 等策略覆盖(settings 表)
|
||||
}
|
||||
|
||||
// NewUserService 创建用户服务。
|
||||
@@ -35,6 +36,20 @@ func NewUserService(db *gorm.DB, sys system.Manager, cfg *config.Config) *UserSe
|
||||
return &UserService{db: db, sys: sys, cfg: cfg}
|
||||
}
|
||||
|
||||
// WithSettings 注入设置服务(settings 表覆盖策略默认值,PLAN F7)。nil 安全。
|
||||
func (s *UserService) WithSettings(settings *SettingService) *UserService {
|
||||
s.settings = settings
|
||||
return s
|
||||
}
|
||||
|
||||
// effectiveDefaultTTL 新账号默认有效期:settings 覆盖优先,否则用配置默认。
|
||||
func (s *UserService) effectiveDefaultTTL(ctx context.Context) time.Duration {
|
||||
if s.settings != nil {
|
||||
return s.settings.DefaultTTL(ctx)
|
||||
}
|
||||
return s.cfg.Policy.DefaultTTL
|
||||
}
|
||||
|
||||
// GetByUsername 按用户名查询外部用户(含或不含 ext_ 前缀均可)。
|
||||
func (s *UserService) GetByUsername(ctx context.Context, username string) (*model.User, error) {
|
||||
full := normalizeName(username, s.cfg.System.UserPrefix)
|
||||
@@ -118,7 +133,7 @@ func (s *UserService) Create(ctx context.Context, username, email, supervisor, p
|
||||
return nil, ErrUserExists
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = s.cfg.Policy.DefaultTTL
|
||||
ttl = s.effectiveDefaultTTL(ctx)
|
||||
}
|
||||
expireAt := time.Now().Add(ttl)
|
||||
u := &model.User{
|
||||
@@ -222,33 +237,37 @@ func (s *UserService) Enable(ctx context.Context, id uint) error {
|
||||
return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Update("status", model.UserStatusActive).Error
|
||||
}
|
||||
|
||||
// Extend 延长有效期:重设 expire_at(days<=0 用配置默认 TTL)。
|
||||
// Extend 延长有效期:重设 expire_at(days<=0 用默认 TTL,settings 覆盖优先)。
|
||||
// 已过期用户在回收期内可经此恢复(PLAN §2.2),恢复后同步密钥。
|
||||
func (s *UserService) Extend(ctx context.Context, id uint, days int) error {
|
||||
// 返回新的过期时间,供 handler 响应。
|
||||
func (s *UserService) Extend(ctx context.Context, id uint, days int) (time.Time, error) {
|
||||
u, err := s.GetByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
return time.Time{}, err
|
||||
}
|
||||
ttl := time.Duration(days) * 24 * time.Hour
|
||||
if days <= 0 {
|
||||
ttl = s.cfg.Policy.DefaultTTL
|
||||
ttl = s.effectiveDefaultTTL(ctx)
|
||||
}
|
||||
newExpire := time.Now().Add(ttl)
|
||||
updates := map[string]any{"expire_at": newExpire}
|
||||
if u.Status == model.UserStatusExpired {
|
||||
if !s.systemAccountOK(ctx, u.Username) {
|
||||
return ErrSystemAccountMissing
|
||||
return time.Time{}, ErrSystemAccountMissing
|
||||
}
|
||||
keys, err := activeUserKeys(s.db, ctx, u.ID, 0)
|
||||
if err != nil {
|
||||
return err
|
||||
return time.Time{}, err
|
||||
}
|
||||
if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, keys); err != nil {
|
||||
return err
|
||||
return time.Time{}, err
|
||||
}
|
||||
updates["status"] = model.UserStatusActive
|
||||
}
|
||||
return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
|
||||
return time.Time{}, err
|
||||
}
|
||||
return newExpire, nil
|
||||
}
|
||||
|
||||
// Delete 删除并回收用户:删除系统账号(userdel -r)+ 家目录 + 密钥记录,
|
||||
|
||||
@@ -182,7 +182,7 @@ func TestUserServiceLifecycle(t *testing.T) {
|
||||
}
|
||||
|
||||
// 延期
|
||||
if err := svc.Extend(ctx, u.ID, 30); err != nil {
|
||||
if _, err := svc.Extend(ctx, u.ID, 30); err != nil {
|
||||
t.Fatalf("extend: %v", err)
|
||||
}
|
||||
|
||||
@@ -219,7 +219,7 @@ func TestUserServiceEnableExpired(t *testing.T) {
|
||||
t.Fatalf("enable expired err = %v, want ErrUserExpired", err)
|
||||
}
|
||||
// 延期可恢复
|
||||
if err := svc.Extend(ctx, u.ID, 30); err != nil {
|
||||
if _, err := svc.Extend(ctx, u.ID, 30); err != nil {
|
||||
t.Fatalf("extend expired: %v", err)
|
||||
}
|
||||
if u2, _ := svc.GetByID(ctx, u.ID); u2.Status != model.UserStatusActive {
|
||||
|
||||
Reference in New Issue
Block a user