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
+153 -10
View File
@@ -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 BOMExcel 直接打开不乱码。
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 BOMExcel 识别中文
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
}
+1 -1
View File
@@ -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{}
+132
View File
@@ -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_keysSSH 立即失效)+ 状态置 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)
}
}
+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)
}
}
+129
View File
@@ -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
View File
@@ -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_atdays<=0 用配置默认 TTL)。
// Extend 延长有效期:重设 expire_atdays<=0 用默认 TTLsettings 覆盖优先)。
// 已过期用户在回收期内可经此恢复(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)+ 家目录 + 密钥记录,
+2 -2
View File
@@ -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 {