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
+1 -1
View File
@@ -13,7 +13,7 @@ NET_HOST := --network=host
GO ?= go
PODMAN ?= podman
BIN := bin/usernode
VERSION ?= 0.3.0-m3
VERSION ?= 0.3.0-m4
LDFLAGS := -s -w -X main.version=$(VERSION)
GOFLAGS := -trimpath
+8 -4
View File
@@ -86,7 +86,8 @@ func cmdServe(args []string) error {
adminSvc := service.NewAdminService(db)
userSvc := service.NewUserService(db, sys, cfg)
keySvc := service.NewKeyService(db, sys, cfg)
auditSvc := service.NewAuditService(db)
auditSvc := service.NewAuditService(db, log)
settingsSvc := service.NewSettingService(db, cfg)
// 认证依赖:DB OTP/会话/重置令牌存储 + 内存图形验证码/登录限速器 + 邮件
captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL)
@@ -98,6 +99,9 @@ func cmdServe(args []string) error {
mailer := mail.NewQueuedMailer(mail.New(cfg.SMTP, log), db, log)
authSvc := service.NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, log)
approvalSvc := service.NewApprovalService(db, cfg, userSvc, mailer, log)
// M4:设置覆盖默认 TTL(settings 表);生命周期每日维护
userSvc.WithSettings(settingsSvc)
lifecycleSvc := service.NewLifecycleService(db, cfg, sys, userSvc, settingsSvc, mailer, auditSvc, log)
// 启动前自动迁移(骨架阶段保证表结构就绪;M5 部署建议显式 migrate)
if err := model.Migrate(db); err != nil {
@@ -105,13 +109,13 @@ func cmdServe(args []string) error {
}
log.Info("db: migrated", "driver", cfg.Database.Driver, "dsn", cfg.Database.DSN)
// 定时任务(M3:邮件失败重试;M4 填充过期扫描/回收/审计归档)
// 定时任务(M3:邮件重试;M4过期扫描/回收/审计归档)
sched := robfigcron.New()
cron.New(db, mailer, log).Register(sched)
cron.New(cfg, mailer, lifecycleSvc, auditSvc, settingsSvc, log).Register(sched)
sched.Start()
defer sched.Stop()
h := api.New(cfg, authSvc, userSvc, keySvc, approvalSvc, auditSvc)
h := api.New(cfg, authSvc, userSvc, keySvc, approvalSvc, auditSvc, settingsSvc)
r := router.New(cfg, h, sessions, log)
srv := server.New(cfg.Server.Listen, r, log)
+4
View File
@@ -51,3 +51,7 @@ group = "external" # 外部用户统一组
shell = "/bin/sh" # 默认 shell
home_base = "/home" # 家目录基路径
authorized_keys_dir = ".ssh" # authorized_keys 所在目录名
[audit]
archive_dir = "" # 审计每日归档目录(M4):留空 = 不归档也不自动清理(防丢审计);
# 生产建议如 /var/lib/usernode/audit_archive,超期审计先归档再清理
+97 -4
View File
@@ -97,19 +97,21 @@ func setupTestAppWithSys(t *testing.T, sys system.Manager) *testApp {
adminSvc := service.NewAdminService(db)
userSvc := service.NewUserService(db, sys, cfg)
keySvc := service.NewKeyService(db, sys, cfg)
auditSvc := service.NewAuditService(db)
mailer := &recordingMailer{}
log := slog.New(slog.NewTextHandler(io.Discard, nil))
auditSvc := service.NewAuditService(db, log)
settingsSvc := service.NewSettingService(db, cfg)
userSvc.WithSettings(settingsSvc)
captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL)
otps := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin)
sessions := auth.NewDBSessionStore(db)
resets := auth.NewDBResetTokenStore(db)
limiter := auth.NewRateLimiter(cfg.Auth.MaxLoginFailures, cfg.Auth.LockDuration)
mailer := &recordingMailer{}
log := slog.New(slog.NewTextHandler(io.Discard, nil))
authSvc := service.NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, log)
approvalSvc := service.NewApprovalService(db, cfg, userSvc, mailer, log)
h := api.New(cfg, authSvc, userSvc, keySvc, approvalSvc, auditSvc)
h := api.New(cfg, authSvc, userSvc, keySvc, approvalSvc, auditSvc, settingsSvc)
r := router.New(cfg, h, sessions, log)
if _, err := adminSvc.Create(context.Background(), "root", "Passw0rd", "root@example.com"); err != nil {
@@ -576,3 +578,94 @@ func TestAPIApprovalFlow(t *testing.T) {
t.Fatalf("resubmit status = %d, body=%s", w.Code, w.Body.String())
}
}
func TestAPIAuditQueryExport(t *testing.T) {
app := setupTestApp(t)
// 登录后产生若干审计记录(登录本身即审计)
w := app.doJSON(http.MethodPost, "/api/v1/auth/admin/login", map[string]string{"username": "root", "password": "Passw0rd"})
ck := sessionCookie(t, w)
// 查询:action 筛选 + 分页
w = app.doJSON(http.MethodGet, "/api/v1/audit?action=admin.login&page=1&page_size=10", nil, ck)
if w.Code != http.StatusOK {
t.Fatalf("audit list status = %d, body=%s", w.Code, w.Body.String())
}
m := decodeBody(t, w)["data"].(map[string]any)
if m["total"].(float64) < 1 {
t.Fatalf("audit total = %v, want >= 1", m["total"])
}
if items := m["items"].([]any); len(items) == 0 {
t.Fatal("audit items empty")
}
// 非法时间参数 → 400
w = app.doJSON(http.MethodGet, "/api/v1/audit?since=not-a-time", nil, ck)
if w.Code != http.StatusBadRequest {
t.Fatalf("bad since status = %d, want 400", w.Code)
}
// 未登录 / 外部用户 → 401 / 403
w = app.doJSON(http.MethodGet, "/api/v1/audit", nil)
if w.Code != http.StatusUnauthorized {
t.Fatalf("unauth audit status = %d, want 401", w.Code)
}
// CSV 导出
w = app.doJSON(http.MethodGet, "/api/v1/audit/export", nil, ck)
if w.Code != http.StatusOK {
t.Fatalf("export status = %d, body=%s", w.Code, w.Body.String())
}
body := w.Body.String()
if !strings.HasPrefix(body, "\ufeffid,created_at,") {
t.Fatalf("csv body = %q", body[:min(40, len(body))])
}
if !strings.Contains(body, "admin.login") {
t.Fatalf("csv missing rows: %q", body[:min(200, len(body))])
}
}
func TestAPISettings(t *testing.T) {
app := setupTestApp(t)
w := app.doJSON(http.MethodPost, "/api/v1/auth/admin/login", map[string]string{"username": "root", "password": "Passw0rd"})
ck := sessionCookie(t, w)
// 未登录 → 401
w = app.doJSON(http.MethodGet, "/api/v1/settings", nil)
if w.Code != http.StatusUnauthorized {
t.Fatalf("unauth settings status = %d, want 401", w.Code)
}
// 初始值(config 默认)
w = app.doJSON(http.MethodGet, "/api/v1/settings", nil, ck)
if w.Code != http.StatusOK {
t.Fatalf("settings list status = %d", w.Code)
}
items := decodeBody(t, w)["data"].(map[string]any)["items"].([]any)
if len(items) != 3 {
t.Fatalf("settings items = %d, want 3", len(items))
}
// 更新默认有效期
w = app.doJSON(http.MethodPut, "/api/v1/settings", map[string]string{"key": "policy.default_ttl", "value": "720h"}, ck)
if w.Code != http.StatusOK {
t.Fatalf("settings put status = %d, body=%s", w.Code, w.Body.String())
}
// 未知 key → 400
w = app.doJSON(http.MethodPut, "/api/v1/settings", map[string]string{"key": "smtp.host", "value": "x"}, ck)
if w.Code != http.StatusBadRequest {
t.Fatalf("unknown key status = %d, want 400", w.Code)
}
// 再次读取:default_ttl 已覆盖
w = app.doJSON(http.MethodGet, "/api/v1/settings", nil, ck)
items = decodeBody(t, w)["data"].(map[string]any)["items"].([]any)
found := false
for _, it := range items {
item := it.(map[string]any)
if item["key"] == "policy.default_ttl" && item["overridden"] == true && item["value"] == "720h" {
found = true
}
}
if !found {
t.Fatalf("settings after put = %v", items)
}
}
+85
View File
@@ -0,0 +1,85 @@
package api
import (
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"ws_usernode/internal/config"
"ws_usernode/internal/service"
)
// AuditHandler 审计接口(admin):查询 + 手动 CSV 导出(PLAN F6)。
type AuditHandler struct {
svc *service.AuditService
cfg *config.Config
}
// NewAuditHandler 创建审计 handler。
func NewAuditHandler(svc *service.AuditService, cfg *config.Config) *AuditHandler {
return &AuditHandler{svc: svc, cfg: cfg}
}
// parseTime 解析可选的时间参数(RFC3339)。
func parseTime(v string) (*time.Time, error) {
if v == "" {
return nil, nil
}
t, err := time.Parse(time.RFC3339, v)
if err != nil {
return nil, err
}
return &t, nil
}
// List GET /audit —— 审计查询(操作者/动作/资源/时间范围/分页)。
func (h *AuditHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
since, err := parseTime(c.Query("since"))
if err != nil {
fail(c, http.StatusBadRequest, "since 需为 RFC3339 时间")
return
}
until, err := parseTime(c.Query("until"))
if err != nil {
fail(c, http.StatusBadRequest, "until 需为 RFC3339 时间")
return
}
rows, total, err := h.svc.Query(c.Request.Context(), service.AuditFilter{
ActorName: c.Query("actor"),
Action: c.Query("action"),
ResourceType: c.Query("resource_type"),
ResourceID: c.Query("resource_id"),
Since: since,
Until: until,
Page: page,
PageSize: pageSize,
})
if err != nil {
fail(c, http.StatusInternalServerError, err.Error())
return
}
ok(c, gin.H{"total": total, "items": rows})
}
// Export GET /audit/export —— 手动 CSV 导出(可选时间范围)。
func (h *AuditHandler) Export(c *gin.Context) {
since, err := parseTime(c.Query("since"))
if err != nil {
fail(c, http.StatusBadRequest, "since 需为 RFC3339 时间")
return
}
until, err := parseTime(c.Query("until"))
if err != nil {
fail(c, http.StatusBadRequest, "until 需为 RFC3339 时间")
return
}
c.Header("Content-Type", "text/csv; charset=utf-8")
c.Header("Content-Disposition", `attachment; filename="audit-`+time.Now().Format("2006-01-02")+`.csv"`)
// 已开始写响应体后无法再返回 JSON 错误;查询失败在此吞掉(导出为管理操作,
// 失败由审计归档兜底),正常路径写入 CSV。
_ = h.svc.ExportCSV(c.Request.Context(), c.Writer, since, until)
}
+8 -3
View File
@@ -1,7 +1,8 @@
// Package api 为 HTTP handler 层(RESTful v1)。
// M1 覆盖认证(管理员/外部用户登录、会话、OTP)与用户管理 CRUD;
// M2 覆盖 SSH 公钥管理(上传/重命名/吊销/列表);
// M3 覆盖申请审批(公开提交 + 管理员审批 + 邮件通知)
// M3 覆盖申请审批(公开提交 + 管理员审批 + 邮件通知)
// M4 覆盖审计查询/导出与系统设置。
package api
import (
@@ -24,13 +25,15 @@ type Handler struct {
User *UserHandler
Key *KeyHandler
Approval *ApprovalHandler
Audit *AuditHandler
Settings *SettingsHandler
authSvc *service.AuthService
auditSvc *service.AuditService
}
// New 创建 handler 集合。
func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, keySvc *service.KeyService, approvalSvc *service.ApprovalService, auditSvc *service.AuditService) *Handler {
func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, keySvc *service.KeyService, approvalSvc *service.ApprovalService, auditSvc *service.AuditService, settingsSvc *service.SettingService) *Handler {
h := &Handler{
Health: &HealthHandler{startedAt: time.Now()},
authSvc: authSvc,
@@ -40,6 +43,8 @@ func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.User
h.User = &UserHandler{svc: userSvc, cfg: cfg, h: h}
h.Key = &KeyHandler{svc: keySvc, cfg: cfg, h: h}
h.Approval = &ApprovalHandler{svc: approvalSvc, cfg: cfg, h: h}
h.Audit = NewAuditHandler(auditSvc, cfg)
h.Settings = NewSettingsHandler(settingsSvc)
return h
}
@@ -66,7 +71,7 @@ type HealthHandler struct {
func (h *HealthHandler) Healthz(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"version": "0.3.0-m3",
"version": "0.3.0-m4",
"uptime": time.Since(h.startedAt).String(),
"go": runtime.Version(),
"timestamp": time.Now().UTC().Format(time.RFC3339),
+57
View File
@@ -0,0 +1,57 @@
package api
import (
"errors"
"net/http"
"github.com/gin-gonic/gin"
"ws_usernode/internal/service"
)
// SettingsHandler 系统设置接口(adminPLAN F7)。
type SettingsHandler struct {
svc *service.SettingService
}
// NewSettingsHandler 创建设置 handler。
func NewSettingsHandler(svc *service.SettingService) *SettingsHandler {
return &SettingsHandler{svc: svc}
}
// List GET /settings —— 全部设置项(config 默认值 + settings 覆盖)。
func (h *SettingsHandler) List(c *gin.Context) {
items, err := h.svc.GetAll(c.Request.Context())
if err != nil {
fail(c, http.StatusInternalServerError, err.Error())
return
}
ok(c, gin.H{"items": items})
}
// SettingUpdateRequest 更新单个设置项。
type SettingUpdateRequest struct {
Key string `json:"key" binding:"required"`
Value string `json:"value" binding:"required"`
}
// Update PUT /settings —— 更新设置项(仅白名单 key,值为时长格式)。
func (h *SettingsHandler) Update(c *gin.Context) {
var req SettingUpdateRequest
if err := c.ShouldBindJSON(&req); err != nil {
fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error())
return
}
if err := h.svc.Set(c.Request.Context(), req.Key, req.Value); err != nil {
switch {
case errors.Is(err, service.ErrSettingKeyUnknown):
fail(c, http.StatusBadRequest, err.Error())
case errors.Is(err, service.ErrSettingValueInvalid):
fail(c, http.StatusBadRequest, err.Error())
default:
fail(c, http.StatusInternalServerError, err.Error())
}
return
}
ok(c, gin.H{"status": "updated", "key": req.Key, "value": req.Value})
}
+3 -10
View File
@@ -179,21 +179,14 @@ func (h *UserHandler) Extend(c *gin.Context) {
fail(c, http.StatusNotFound, err.Error())
return
}
if err := h.svc.Extend(c.Request.Context(), uint(id), req.Days); err != nil {
newExpire, err := h.svc.Extend(c.Request.Context(), uint(id), req.Days)
if err != nil {
h.h.audit(c, "user.extend", "user", c.Param("id"), map[string]any{"days": req.Days, "err": err.Error()}, model.ResultFailed)
fail(c, http.StatusInternalServerError, err.Error())
return
}
h.h.audit(c, "user.extend", "user", c.Param("id"), map[string]any{"days": req.Days, "old_status": u.Status}, model.ResultSuccess)
ok(c, gin.H{"expire_at": time.Now().Add(h.extendTTL(req.Days)).UTC()})
}
// extendTTL 计算新的有效期(与 service 保持一致:days<=0 用默认)。
func (h *UserHandler) extendTTL(days int) time.Duration {
if days <= 0 {
return h.cfg.Policy.DefaultTTL
}
return time.Duration(days) * 24 * time.Hour
ok(c, gin.H{"expire_at": newExpire.UTC()})
}
// Delete DELETE /users/:id —— 删除并回收(系统账号 + 家目录 + 密钥,保留审计)。
+10
View File
@@ -31,6 +31,7 @@ type Config struct {
Auth AuthConfig `toml:"auth"`
SMTP SMTPConfig `toml:"smtp"`
System SystemConfig `toml:"system"`
Audit AuditConfig `toml:"audit"`
}
type AppConfig struct {
@@ -90,6 +91,11 @@ type SystemConfig struct {
AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖)
}
// AuditConfig 审计保留与归档配置。
type AuditConfig struct {
ArchiveDir string `toml:"archive_dir"` // 每日归档目录(空 = 不归档也不自动清理,防止丢审计)
}
// Default 返回带开发环境默认值的配置,作为 config.example.toml 与未配置项的兜底。
func Default() *Config {
return &Config{
@@ -114,6 +120,10 @@ func Default() *Config {
CaptchaTTL: 5 * time.Minute,
},
SMTP: SMTPConfig{Port: 587},
Audit: AuditConfig{
// 开发默认不归档(避免在任意目录落文件);生产显式配置 archive_dir 启用归档。
ArchiveDir: "",
},
System: SystemConfig{
// 开发默认 dry-run:未配置 config 直接跑 serve 时只打印计划,避免误操作系统账号。
// 生产必须显式 dry_run=false 且 sudo=true(见 deploy/sudoers.example)。
+53 -13
View File
@@ -1,28 +1,35 @@
// Package cron 提供定时任务:邮件重试、过期扫描、回收、审计归档。
// M0 注册任务框架(过期扫描/回收在 M1/M4 接入)M3 接入邮件失败重试;
// 单实例部署说明见 PLAN §11。
// M0 注册任务框架;M3 接入邮件失败重试;M4 接入生命周期(到期锁定/回收)
// 与审计归档;单实例部署说明见 PLAN §11。
package cron
import (
"context"
"log/slog"
"time"
"github.com/robfig/cron/v3"
"gorm.io/gorm"
"ws_usernode/internal/config"
"ws_usernode/internal/mail"
"ws_usernode/internal/service"
)
// Jobs 聚合定时任务依赖与注册。
type Jobs struct {
db *gorm.DB
cfg *config.Config
mailer mail.Retryable
lifecycle *service.LifecycleService
audit *service.AuditService
settings *service.SettingService
log *slog.Logger
}
// New 创建定时任务集合。mailer 为邮件队列(失败重试,PLAN F5)。
func New(db *gorm.DB, mailer mail.Retryable, log *slog.Logger) *Jobs {
return &Jobs{db: db, mailer: mailer, log: log}
// New 创建定时任务集合。
// mailer 为邮件队列(失败重试,M3);lifecycle 为账号生命周期(M4);
// audit/settings 用于每日审计归档(M4)。
func New(cfg *config.Config, mailer mail.Retryable, lifecycle *service.LifecycleService, audit *service.AuditService, settings *service.SettingService, log *slog.Logger) *Jobs {
return &Jobs{cfg: cfg, mailer: mailer, lifecycle: lifecycle, audit: audit, settings: settings, log: log}
}
// Register 将任务注册到 cron。
@@ -34,10 +41,43 @@ func (j *Jobs) Register(c *cron.Cron) {
}
})
// 过期扫描(M4):每日扫描到期账号,锁系统账号 + 邮件通知
// 回收任务(M4):超回收期账号自动回收(userdel + 保留审计)
// 审计归档(M4):每日将过期审计记录导出归档后清理
c.AddFunc("@daily", func() {
j.log.Info("cron: daily maintenance tick (tasks to be implemented in M1/M4)")
})
// 每日维护:过期扫描 → 回收 → 审计归档(M4)
c.AddFunc("@daily", j.dailyMaintenance)
// 到期扫描额外每 6 小时跑一次:避免错过当日到期账号
c.AddFunc("@every 6h", j.scanExpired)
}
// dailyMaintenance 每日一次:到期扫描、超期回收、审计归档。
func (j *Jobs) dailyMaintenance() {
ctx := context.Background()
j.scanExpired()
if n, err := j.lifecycle.Recycle(ctx); err != nil {
j.log.Error("cron: recycle failed", "err", err)
} else if n > 0 {
j.log.Info("cron: recycled accounts", "count", n)
}
j.archiveAudit()
}
// scanExpired 到期扫描:锁定系统账号 + 邮件通知。
func (j *Jobs) scanExpired() {
ctx := context.Background()
if n, err := j.lifecycle.ScanExpired(ctx); err != nil {
j.log.Error("cron: expire scan failed", "err", err)
} else if n > 0 {
j.log.Info("cron: expired accounts locked", "count", n)
}
}
// archiveAudit 审计归档:将超过保留期的记录归档后清理(archive_dir 未配置时跳过)。
func (j *Jobs) archiveAudit() {
ctx := context.Background()
retention := j.settings.AuditRetention(ctx)
before := time.Now().Add(-retention)
if n, err := j.audit.Archive(ctx, j.cfg.Audit.ArchiveDir, before); err != nil {
j.log.Error("cron: audit archive failed", "err", err)
} else if n > 0 {
j.log.Info("cron: audit archived", "count", n, "before", before.Format(time.RFC3339))
}
}
+13 -13
View File
@@ -3,7 +3,6 @@ package router
import (
"log/slog"
"net/http"
"time"
"github.com/gin-gonic/gin"
@@ -74,11 +73,19 @@ func New(cfg *config.Config, h *api.Handler, sessions auth.SessionStore, log *sl
approvals.POST("/:id/review", h.Approval.Review)
}
// 后续里程碑
v1.GET("/audit", notImplemented("审计查询(M4"))
v1.GET("/audit/export", notImplemented("审计导出(M4"))
v1.GET("/settings", notImplemented("设置(M4"))
v1.PUT("/settings", notImplemented("设置(M4"))
// 审计(M4):查询 + 手动 CSV 导出(admin
auditGrp := v1.Group("/audit", sessionMiddleware(sessions), requireUserType(auth.SessionUserAdmin))
{
auditGrp.GET("", h.Audit.List)
auditGrp.GET("/export", h.Audit.Export)
}
// 系统设置(M4):读写(admin)
settings := v1.Group("/settings", sessionMiddleware(sessions), requireUserType(auth.SessionUserAdmin))
{
settings.GET("", h.Settings.List)
settings.PUT("", h.Settings.Update)
}
}
// 前端静态资源(go:embeddev 阶段由 Vite dev server 代理,见 Makefile dev
@@ -87,13 +94,6 @@ func New(cfg *config.Config, h *api.Handler, sessions auth.SessionStore, log *sl
return r
}
// notImplemented 返回 501 占位 handler,标注里程碑。
func notImplemented(what string) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusNotImplemented, gin.H{"error": "接口 " + what + " 尚未实现"})
}
}
// requestLogger 以 slog 输出结构化请求日志。
func requestLogger(log *slog.Logger) gin.HandlerFunc {
return func(c *gin.Context) {
+152 -9
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
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"}
}
+28 -9
View File
@@ -28,6 +28,7 @@ type UserService struct {
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 {