From c0e2ff975a1fc5a097e5270de56bae0b519fe8ab Mon Sep 17 00:00:00 2001 From: CaoWangrenbo Date: Sun, 30 Aug 2026 10:17:22 +0800 Subject: [PATCH] =?UTF-8?q?feat(M4):=20=E7=94=9F=E5=91=BD=E5=91=A8?= =?UTF-8?q?=E6=9C=9F=20+=20=E5=AE=A1=E8=AE=A1=20=E2=80=94=20=E5=88=B0?= =?UTF-8?q?=E6=9C=9F=E9=94=81=E5=AE=9A/=E5=9B=9E=E6=94=B6=20cron=E3=80=81?= =?UTF-8?q?=E5=AE=A1=E8=AE=A1=E6=9F=A5=E8=AF=A2/CSV=20=E5=AF=BC=E5=87=BA/?= =?UTF-8?q?=E6=AF=8F=E6=97=A5=E5=BD=92=E6=A1=A3=E3=80=81settings=20?= =?UTF-8?q?=E5=8A=A8=E6=80=81=E7=AD=96=E7=95=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Makefile | 2 +- cmd/usernode/app.go | 12 +- config.example.toml | 4 + internal/api/api_test.go | 101 +++++++++++- internal/api/audit.go | 85 ++++++++++ internal/api/handler.go | 11 +- internal/api/settings.go | 57 +++++++ internal/api/user.go | 13 +- internal/config/config.go | 10 ++ internal/cron/cron.go | 70 ++++++-- internal/router/router.go | 26 +-- internal/service/audit.go | 163 ++++++++++++++++-- internal/service/auth_test.go | 2 +- internal/service/lifecycle.go | 132 +++++++++++++++ internal/service/m4_test.go | 299 ++++++++++++++++++++++++++++++++++ internal/service/settings.go | 129 +++++++++++++++ internal/service/user.go | 43 +++-- internal/service/user_test.go | 4 +- 18 files changed, 1088 insertions(+), 75 deletions(-) create mode 100644 internal/api/audit.go create mode 100644 internal/api/settings.go create mode 100644 internal/service/lifecycle.go create mode 100644 internal/service/m4_test.go create mode 100644 internal/service/settings.go diff --git a/Makefile b/Makefile index b70ce30..0ce364a 100644 --- a/Makefile +++ b/Makefile @@ -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 diff --git a/cmd/usernode/app.go b/cmd/usernode/app.go index a3ed6e2..db9989d 100644 --- a/cmd/usernode/app.go +++ b/cmd/usernode/app.go @@ -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) diff --git a/config.example.toml b/config.example.toml index f64a12c..ef30935 100644 --- a/config.example.toml +++ b/config.example.toml @@ -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,超期审计先归档再清理 diff --git a/internal/api/api_test.go b/internal/api/api_test.go index 45736d3..97fb6a8 100644 --- a/internal/api/api_test.go +++ b/internal/api/api_test.go @@ -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) + } +} diff --git a/internal/api/audit.go b/internal/api/audit.go new file mode 100644 index 0000000..83b5095 --- /dev/null +++ b/internal/api/audit.go @@ -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) +} diff --git a/internal/api/handler.go b/internal/api/handler.go index d510de0..835c49c 100644 --- a/internal/api/handler.go +++ b/internal/api/handler.go @@ -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), diff --git a/internal/api/settings.go b/internal/api/settings.go new file mode 100644 index 0000000..ed3d02a --- /dev/null +++ b/internal/api/settings.go @@ -0,0 +1,57 @@ +package api + +import ( + "errors" + "net/http" + + "github.com/gin-gonic/gin" + + "ws_usernode/internal/service" +) + +// SettingsHandler 系统设置接口(admin,PLAN 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}) +} diff --git a/internal/api/user.go b/internal/api/user.go index 5a72628..fa99d8e 100644 --- a/internal/api/user.go +++ b/internal/api/user.go @@ -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 —— 删除并回收(系统账号 + 家目录 + 密钥,保留审计)。 diff --git a/internal/config/config.go b/internal/config/config.go index 47edf85..ef8c6b3 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -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)。 diff --git a/internal/cron/cron.go b/internal/cron/cron.go index 41b33d8..2b4421c 100644 --- a/internal/cron/cron.go +++ b/internal/cron/cron.go @@ -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 - mailer mail.Retryable - log *slog.Logger + 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)) + } } diff --git a/internal/router/router.go b/internal/router/router.go index c71a5a8..193c1fd 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -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:embed;dev 阶段由 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) { diff --git a/internal/service/audit.go b/internal/service/audit.go index 28eb1f4..e1ec900 100644 --- a/internal/service/audit.go +++ b/internal/service/audit.go @@ -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 } diff --git a/internal/service/auth_test.go b/internal/service/auth_test.go index b25c8bd..37ca918 100644 --- a/internal/service/auth_test.go +++ b/internal/service/auth_test.go @@ -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{} diff --git a/internal/service/lifecycle.go b/internal/service/lifecycle.go new file mode 100644 index 0000000..57917e1 --- /dev/null +++ b/internal/service/lifecycle.go @@ -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) + } +} diff --git a/internal/service/m4_test.go b/internal/service/m4_test.go new file mode 100644 index 0000000..d1a2000 --- /dev/null +++ b/internal/service/m4_test.go @@ -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) + } +} diff --git a/internal/service/settings.go b/internal/service/settings.go new file mode 100644 index 0000000..be71c99 --- /dev/null +++ b/internal/service/settings.go @@ -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"} +} diff --git a/internal/service/user.go b/internal/service/user.go index f48d2fa..d356aea 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -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)+ 家目录 + 密钥记录, diff --git a/internal/service/user_test.go b/internal/service/user_test.go index 68f7321..34fcd33 100644 --- a/internal/service/user_test.go +++ b/internal/service/user_test.go @@ -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 {