From 630d240dc04f60e26c615f66b8091690adfcfaf3 Mon Sep 17 00:00:00 2001 From: CaoWangrenbo Date: Sat, 29 Aug 2026 23:40:20 +0800 Subject: [PATCH] =?UTF-8?q?feat(M1):=20=E8=AE=A4=E8=AF=81=E4=B8=8E?= =?UTF-8?q?=E7=94=A8=E6=88=B7=E7=AE=A1=E7=90=86=20=E2=80=94=20=E5=8F=8C?= =?UTF-8?q?=E9=80=9A=E9=81=93=E7=99=BB=E5=BD=95=E3=80=81cookie=20=E4=BC=9A?= =?UTF-8?q?=E8=AF=9D=E3=80=81=E7=94=A8=E6=88=B7=20CRUD=20=E4=B8=8E?= =?UTF-8?q?=E7=9C=9F=E5=AE=9E=E7=B3=BB=E7=BB=9F=E8=B4=A6=E5=8F=B7=E5=AF=B9?= =?UTF-8?q?=E6=8E=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 认证: - 图形验证码 GET /auth/captcha(内置 PNG 渲染,零第三方依赖) - 外部用户 OTP 双通道:DB 存储(otp_codes)使邮件与 CLI 共用同一验证码/冷却/失败限速 - 管理员 bcrypt 登录 + 连续失败限速锁定;admin/forgot + admin/reset 邮件重置(SMTP 或日志) - cookie 会话(HttpOnly/SameSite)、me/logout、admin/user 鉴权中间件 用户管理(admin): - CRUD + disable/enable/extend/delete,对接 system 层真实 useradd/usermod/userdel/passwd - system 层三执行模式:dry-run(默认,安全)/ direct(容器/测试用户)/ sudo(生产 sudoers 白名单) - Exists 系统账号一致性检查;deploy/sudoers.example 白名单模板 - 关键操作接入 append-only 审计 其他: - CLI user otp 改 DB store,与邮件通道真正对齐 - 容器镜像补 shadow(alpine 无 useradd);Makefile VERSION 0.2.0-m1 - 测试:auth/service 单测 + api httptest 集成 + 容器内真实系统账号端到端验证 --- Makefile | 2 +- cmd/usernode/app.go | 17 +- cmd/usernode/user.go | 35 ++-- config.example.toml | 9 +- deploy/Containerfile | 3 +- deploy/sudoers.example | 23 +++ internal/api/admin.go | 53 ------ internal/api/api_test.go | 326 +++++++++++++++++++++++++++++++++ internal/api/auth.go | 195 ++++++++++++++++++++ internal/api/cookie.go | 40 ++++ internal/api/handler.go | 51 +++++- internal/api/user.go | 189 +++++++++++++++++-- internal/auth/captcha.go | 12 +- internal/auth/captcha_img.go | 131 +++++++++++++ internal/auth/captcha_test.go | 115 ++++++++++++ internal/auth/otp.go | 176 ++++++++++++++++-- internal/auth/otp_test.go | 119 ++++++++++++ internal/auth/ratelimit.go | 68 +++++++ internal/auth/reset.go | 76 ++++++++ internal/config/config.go | 40 ++-- internal/config/config_test.go | 16 ++ internal/mail/mailer.go | 118 ++++++++++++ internal/model/model.go | 32 ++++ internal/router/middleware.go | 46 +++++ internal/router/router.go | 42 +++-- internal/service/admin.go | 44 ++++- internal/service/auth.go | 235 ++++++++++++++++++++++++ internal/service/auth_test.go | 230 +++++++++++++++++++++++ internal/service/user.go | 246 +++++++++++++++++++++++-- internal/service/user_test.go | 242 ++++++++++++++++++++++++ internal/system/system.go | 93 +++++++--- internal/system/system_test.go | 87 +++++++++ 32 files changed, 2923 insertions(+), 188 deletions(-) create mode 100644 deploy/sudoers.example delete mode 100644 internal/api/admin.go create mode 100644 internal/api/api_test.go create mode 100644 internal/api/auth.go create mode 100644 internal/api/cookie.go create mode 100644 internal/auth/captcha_img.go create mode 100644 internal/auth/captcha_test.go create mode 100644 internal/auth/otp_test.go create mode 100644 internal/auth/ratelimit.go create mode 100644 internal/auth/reset.go create mode 100644 internal/mail/mailer.go create mode 100644 internal/router/middleware.go create mode 100644 internal/service/auth.go create mode 100644 internal/service/auth_test.go create mode 100644 internal/service/user_test.go create mode 100644 internal/system/system_test.go diff --git a/Makefile b/Makefile index 35fc3a8..bddd808 100644 --- a/Makefile +++ b/Makefile @@ -13,7 +13,7 @@ NET_HOST := --network=host GO ?= go PODMAN ?= podman BIN := bin/usernode -VERSION ?= 0.1.0-m0 +VERSION ?= 0.2.0-m1 LDFLAGS := -s -w -X main.version=$(VERSION) GOFLAGS := -trimpath diff --git a/cmd/usernode/app.go b/cmd/usernode/app.go index 298a3c8..beda986 100644 --- a/cmd/usernode/app.go +++ b/cmd/usernode/app.go @@ -10,8 +10,10 @@ import ( "gorm.io/gorm" "ws_usernode/internal/api" + "ws_usernode/internal/auth" "ws_usernode/internal/config" "ws_usernode/internal/cron" + "ws_usernode/internal/mail" "ws_usernode/internal/model" "ws_usernode/internal/router" "ws_usernode/internal/server" @@ -82,9 +84,18 @@ func cmdServe(args []string) error { sys := system.New(cfg.System) adminSvc := service.NewAdminService(db) - userSvc := service.NewUserService(db, sys) + userSvc := service.NewUserService(db, sys, cfg) auditSvc := service.NewAuditService(db) + // 认证依赖:DB OTP/会话/重置令牌存储 + 内存图形验证码/登录限速器 + 邮件 + 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 := mail.New(cfg.SMTP, log) + authSvc := service.NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, log) + // 启动前自动迁移(骨架阶段保证表结构就绪;M5 部署建议显式 migrate) if err := model.Migrate(db); err != nil { return fmt.Errorf("数据库迁移: %w", err) @@ -97,8 +108,8 @@ func cmdServe(args []string) error { sched.Start() defer sched.Stop() - h := api.New(adminSvc, userSvc, auditSvc) - r := router.New(cfg, h, log) + h := api.New(cfg, authSvc, userSvc, auditSvc) + r := router.New(cfg, h, sessions, log) srv := server.New(cfg.Server.Listen, r, log) if err := srv.Run(); err != nil { diff --git a/cmd/usernode/user.go b/cmd/usernode/user.go index 166f2ad..6057177 100644 --- a/cmd/usernode/user.go +++ b/cmd/usernode/user.go @@ -25,8 +25,8 @@ func cmdUser(args []string) error { } } -// userOTP 获取外部用户 OTP 验证码,与邮件通道共用同一存储与限速 -// (同一验证码、同一 10 分钟有效期、同一 60s 冷却与失败限速)。 +// userOTP 获取外部用户 OTP 验证码。与邮件通道共用同一 DB 存储与限速: +// 已有有效验证码时直接复用(同一验证码),无则生成(受同一 60s 冷却约束)。 func userOTP(args []string) error { fs := flag.NewFlagSet("user otp", flag.ContinueOnError) cfgPath, debug := commonFlags(fs) @@ -48,27 +48,28 @@ func userOTP(args []string) error { } // 校验用户存在(不存在时返回友好错误,避免暴露账号是否存在的枚举) - if _, err := service.NewUserService(db, system.New(cfg.System)).GetByUsername(context.Background(), name); err != nil { + if _, err := service.NewUserService(db, system.New(cfg.System), cfg).GetByUsername(context.Background(), name); err != nil { return fmt.Errorf("用户不存在或不可用: %w", err) } - // M0 骨架:CLI 独立生成(内存 store 与运行中服务不共享)。 - // 生产对齐(同一验证码/冷却/限速跨通道生效)需 OTP 落 DB,M1 实现 - // auth.OTPStore 的 DB 实现后,CLI 与邮件通道读写同一存储。 - store := auth.NewMemoryOTPStore() - code, err := store.Send(name, cfg.Policy.OTPTTL, cfg.Policy.OTPCooldown) + store := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin) + ctx := context.Background() + code, err := store.Current(ctx, name) if err != nil { - if err == auth.ErrCooldown { - return fmt.Errorf("发送冷却中,请稍后重试(冷却 %s)", cfg.Policy.OTPCooldown) + // 无有效验证码:生成(先到先得,覆盖旧码;冷却期内返回 ErrCooldown) + code, err = store.Send(ctx, name, cfg.Policy.OTPTTL, cfg.Policy.OTPCooldown) + if err != nil { + if err == auth.ErrCooldown { + return fmt.Errorf("发送冷却中,请稍后重试(冷却 %s)", cfg.Policy.OTPCooldown) + } + return err } - return err + log.Info("otp generated", "username", name, "valid_for", cfg.Policy.OTPTTL.String()) + } else { + log.Info("otp reused (与邮件通道同一验证码)", "username", name) } - log.Info("otp generated", - "username", name, - "valid_for", cfg.Policy.OTPTTL.String(), - "expires_at", time.Now().Add(cfg.Policy.OTPTTL).Format(time.RFC3339), - "hint", "与邮件通道为同一验证码,登录后立即失效", - ) fmt.Printf("OTP for %s: %s\n", name, code) + fmt.Printf("有效期至 %s,登录后立即失效;如需重发请等待冷却 %s 或稍后在网页重新请求。\n", + time.Now().Add(cfg.Policy.OTPTTL).Format(time.RFC3339), cfg.Policy.OTPCooldown) return nil } diff --git a/config.example.toml b/config.example.toml index e20de9c..0de141b 100644 --- a/config.example.toml +++ b/config.example.toml @@ -8,6 +8,7 @@ [app] name = "ws_usernode" env = "development" # development / production +base_url = "http://127.0.0.1:8080" # 对外访问地址(邮件重置链接等) [server] listen = "127.0.0.1:8080" # 生产建议 0.0.0.0:8080 并置于反向代理后 @@ -30,6 +31,11 @@ audit_retention = "720h" # 审计保留 30 天,保留前先归档 otp_ttl = "10m" # OTP 验证码有效期 otp_cooldown = "60s" # OTP 发送冷却 +[auth] +max_login_failures = 5 # 管理员登录连续失败阈值,达到后锁定 +lock_duration = "15m" # 锁定持续时间 +captcha_ttl = "5m" # 图形验证码有效期 + [smtp] host = "" # 留空则禁用邮件(OTP 仍可用 CLI 通道获取) port = 587 @@ -38,7 +44,8 @@ password = "" from = "usernode@example.com" [system] -sudo = false # 开发环境 false = dry-run(只打印不执行);生产 true 经 sudo -n 执行 +sudo = false # 生产 true:经 sudo -n 执行 useradd/usermod/userdel/passwd(需 deploy/sudoers) +dry_run = true # 开发演练 true:只打印计划命令不执行;false 且 sudo=false 时直接执行(容器/测试用户验证) user_prefix = "ext_" # 外部用户系统账号统一前缀 group = "external" # 外部用户统一组 shell = "/bin/sh" # 默认 shell diff --git a/deploy/Containerfile b/deploy/Containerfile index 94a853f..c8419b9 100644 --- a/deploy/Containerfile +++ b/deploy/Containerfile @@ -49,7 +49,8 @@ RUN CGO_ENABLED=0 go build -trimpath -ldflags="-s -w" -o /out/usernode ./cmd/use # ---------- 阶段 3:运行镜像 ---------- FROM alpine:3.20 -RUN apk add --no-cache ca-certificates tzdata \ +# shadow 提供 useradd/usermod/userdel/passwd(alpine 默认 busybox 无 useradd) +RUN apk add --no-cache ca-certificates tzdata shadow \ && addgroup -S usernode && adduser -S -G usernode usernode COPY --from=go-build /out/usernode /usr/local/bin/usernode diff --git a/deploy/sudoers.example b/deploy/sudoers.example new file mode 100644 index 0000000..417f6a3 --- /dev/null +++ b/deploy/sudoers.example @@ -0,0 +1,23 @@ +# ws_usernode 节点专有用户 sudoers 白名单(生产部署) +# +# 安装:将本文件复制为 /etc/sudoers.d/usernode 并执行 `visudo -c` 校验。 +# 节点进程以 usernode 用户运行,仅允许以 root 执行固定命令(禁任意 shell), +# 命令参数由程序内强校验(pkg.ValidateSystemAccount 等),见 PLAN §9。 +# +# 注意:以下命令路径基于 Debian/Ubuntu(/usr/sbin)。Alpine 为 /usr/sbin; +# 请按发行版调整,并确保 usernode 用户无 NOPASSWD 的通用提权入口。 + +usernode ALL=(root) NOPASSWD: /usr/sbin/useradd, /usr/sbin/usermod, \ + /usr/sbin/userdel, /usr/bin/passwd + +# 说明: +# - useradd -m -d -s -g external 创建账号 +# - usermod 预留(如 usermod -e 过期),M4 回收期使用 +# - userdel -r 删除账号及家目录 +# - passwd -l / -u 锁定/解锁口令 +# - 不授予 chsh/其他命令的任意执行;若需变更默认 shell 请收紧为固定参数 +# +# 生产禁止 system.sudo=false 的 direct 模式:必须显式配置 +# [system] +# sudo = true +# dry_run = false diff --git a/internal/api/admin.go b/internal/api/admin.go deleted file mode 100644 index 1a6bf90..0000000 --- a/internal/api/admin.go +++ /dev/null @@ -1,53 +0,0 @@ -package api - -import ( - "net/http" - "strings" - - "github.com/gin-gonic/gin" - - "ws_usernode/internal/service" -) - -// AdminHandler 管理员账号相关接口(M1 完成登录;create/reset 走 CLI)。 -type AdminHandler struct { - svc *service.AdminService -} - -// AdminCreateRequest 管理员创建请求。 -type AdminCreateRequest struct { - Username string `json:"username" binding:"required"` - Password string `json:"password" binding:"required"` - Email string `json:"email" binding:"required"` -} - -// Create 创建管理员(仅初始引导用,M1 前可经此接口快速建号)。 -func (h *AdminHandler) Create(c *gin.Context) { - if h.svc == nil { - fail(c, http.StatusNotImplemented, "管理员服务未初始化") - return - } - var req AdminCreateRequest - if err := c.ShouldBindJSON(&req); err != nil { - fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) - return - } - adm, err := h.svc.Create(c.Request.Context(), req.Username, req.Password, req.Email) - if err != nil { - switch { - case strings.Contains(err.Error(), "已存在"): - fail(c, http.StatusConflict, err.Error()) - case strings.Contains(err.Error(), "过弱"): - fail(c, http.StatusBadRequest, err.Error()) - default: - fail(c, http.StatusInternalServerError, err.Error()) - } - return - } - ok(c, gin.H{"id": adm.ID, "username": adm.Username}) -} - -// Me 返回当前管理员(M1 接入会话后启用)。 -func (h *AdminHandler) Me(c *gin.Context) { - fail(c, http.StatusNotImplemented, "会话尚未接入(M1)") -} diff --git a/internal/api/api_test.go b/internal/api/api_test.go new file mode 100644 index 0000000..cb06930 --- /dev/null +++ b/internal/api/api_test.go @@ -0,0 +1,326 @@ +package api_test + +import ( + "bytes" + "context" + "encoding/json" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strconv" + "testing" + + "github.com/gin-gonic/gin" + "gorm.io/gorm" + + "ws_usernode/internal/api" + "ws_usernode/internal/auth" + "ws_usernode/internal/config" + "ws_usernode/internal/model" + "ws_usernode/internal/router" + "ws_usernode/internal/service" + "ws_usernode/internal/system" +) + +// recordingMailer 捕获邮件,用于从重置邮件提取 token 等。 +type recordingMailer struct { + lastTo string + lastSubject string + lastBody string +} + +func (m *recordingMailer) Send(_ context.Context, to, subject, body string) error { + m.lastTo = to + m.lastSubject = subject + m.lastBody = body + return nil +} + +// testApp 完整组装的应用(SQLite 内存库 + dry-run 系统层)。 +type testApp struct { + r http.Handler + db *gorm.DB + captchas auth.CaptchaStore + otps auth.OTPStore + mailer *recordingMailer +} + +func setupTestApp(t *testing.T) *testApp { + t.Helper() + gin.SetMode(gin.TestMode) + db, err := model.Open("sqlite", ":memory:", false) + if err != nil { + t.Fatalf("open db: %v", err) + } + if err := model.Migrate(db); err != nil { + t.Fatalf("migrate: %v", err) + } + cfg := config.Default() + cfg.System.DryRun = true // 集成测试走 dry-run,不触碰真实系统账号 + + sys := system.New(cfg.System) + adminSvc := service.NewAdminService(db) + userSvc := service.NewUserService(db, sys, cfg) + auditSvc := service.NewAuditService(db) + + 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) + + h := api.New(cfg, authSvc, userSvc, auditSvc) + r := router.New(cfg, h, sessions, log) + + if _, err := adminSvc.Create(context.Background(), "root", "Passw0rd", "root@example.com"); err != nil { + t.Fatalf("seed admin: %v", err) + } + return &testApp{r: r, db: db, captchas: captchas, otps: otps, mailer: mailer} +} + +// doJSON 发起 JSON 请求,返回 recorder。 +func (a *testApp) doJSON(method, path string, body any, cookies ...*http.Cookie) *httptest.ResponseRecorder { + var rdr io.Reader + if body != nil { + b, _ := json.Marshal(body) + rdr = bytes.NewReader(b) + } + req := httptest.NewRequest(method, path, rdr) + if body != nil { + req.Header.Set("Content-Type", "application/json") + } + for _, c := range cookies { + req.AddCookie(c) + } + w := httptest.NewRecorder() + a.r.ServeHTTP(w, req) + return w +} + +func decodeBody(t *testing.T, w *httptest.ResponseRecorder) map[string]any { + t.Helper() + var m map[string]any + if err := json.Unmarshal(w.Body.Bytes(), &m); err != nil { + t.Fatalf("decode response %q: %v", w.Body.String(), err) + } + return m +} + +func sessionCookie(t *testing.T, w *httptest.ResponseRecorder) *http.Cookie { + t.Helper() + for _, c := range w.Result().Cookies() { + if c.Name == api.SessionCookieName { + return c + } + } + t.Fatalf("no session cookie in response") + return nil +} + +func TestAPICaptcha(t *testing.T) { + app := setupTestApp(t) + w := app.doJSON(http.MethodGet, "/api/v1/auth/captcha", nil) + if w.Code != http.StatusOK { + t.Fatalf("captcha status = %d, body=%s", w.Code, w.Body.String()) + } + m := decodeBody(t, w) + data := m["data"].(map[string]any) + if data["captcha_id"] == "" || data["image"] == "" { + t.Fatalf("captcha response missing fields: %v", data) + } +} + +func TestAPIAdminLoginRequiresSession(t *testing.T) { + app := setupTestApp(t) + // 未登录访问 /users 应 401 + w := app.doJSON(http.MethodGet, "/api/v1/users", nil) + if w.Code != http.StatusUnauthorized { + t.Fatalf("unauth users status = %d, want 401", w.Code) + } + // 错误密码 401 + w = app.doJSON(http.MethodPost, "/api/v1/auth/admin/login", map[string]string{"username": "root", "password": "bad"}) + if w.Code != http.StatusUnauthorized { + t.Fatalf("bad login status = %d, want 401", w.Code) + } +} + +func TestAPIAdminUserLifecycle(t *testing.T) { + app := setupTestApp(t) + + // 管理员登录 + w := app.doJSON(http.MethodPost, "/api/v1/auth/admin/login", map[string]string{"username": "root", "password": "Passw0rd"}) + if w.Code != http.StatusOK { + t.Fatalf("admin login status = %d, body=%s", w.Code, w.Body.String()) + } + ck := sessionCookie(t, w) + + // me + w = app.doJSON(http.MethodGet, "/api/v1/auth/me", nil, ck) + if w.Code != http.StatusOK { + t.Fatalf("me status = %d", w.Code) + } + if m := decodeBody(t, w); m["data"].(map[string]any)["username"] != "root" { + t.Fatalf("me body = %s", w.Body.String()) + } + + // 创建用户(dry-run 系统层) + w = app.doJSON(http.MethodPost, "/api/v1/users", map[string]any{ + "username": "zhangsan", "email": "zs@example.com", "supervisor": "prof.li", "purpose": "科研", "ttl_days": 90, + }, ck) + if w.Code != http.StatusOK { + t.Fatalf("create user status = %d, body=%s", w.Code, w.Body.String()) + } + created := decodeBody(t, w)["data"].(map[string]any) + id := uint(created["id"].(float64)) + if created["username"] != "ext_zhangsan" { + t.Fatalf("created username = %v", created["username"]) + } + + // 列表 + w = app.doJSON(http.MethodGet, "/api/v1/users?status=active&page=1&page_size=10", nil, ck) + if w.Code != http.StatusOK { + t.Fatalf("list status = %d", w.Code) + } + list := decodeBody(t, w)["data"].(map[string]any) + if list["total"].(float64) != 1 { + t.Fatalf("list total = %v, want 1", list["total"]) + } + + // 详情 + w = app.doJSON(http.MethodGet, "/api/v1/users/"+itoa(id), nil, ck) + if w.Code != http.StatusOK { + t.Fatalf("get status = %d, body=%s", w.Code, w.Body.String()) + } + + // 更新(改邮箱) + email := "zs-new@example.com" + w = app.doJSON(http.MethodPatch, "/api/v1/users/"+itoa(id), map[string]any{"email": email}, ck) + if w.Code != http.StatusOK { + t.Fatalf("update status = %d, body=%s", w.Code, w.Body.String()) + } + + // 禁用 → 启用 → 延期 + for _, action := range []string{"disable", "enable", "extend"} { + body := any(nil) + if action == "extend" { + body = map[string]any{"days": 30} + } + w = app.doJSON(http.MethodPost, "/api/v1/users/"+itoa(id)+"/"+action, body, ck) + if w.Code != http.StatusOK { + t.Fatalf("%s status = %d, body=%s", action, w.Code, w.Body.String()) + } + } + + // 删除 + w = app.doJSON(http.MethodDelete, "/api/v1/users/"+itoa(id), nil, ck) + if w.Code != http.StatusOK { + t.Fatalf("delete status = %d, body=%s", w.Code, w.Body.String()) + } + + // 审计应已写入 + var n int64 + if err := app.db.Model(&model.AuditLog{}).Count(&n).Error; err != nil { + t.Fatalf("audit count: %v", err) + } + if n < 7 { + t.Fatalf("audit entries = %d, want >= 7", n) + } +} + +func TestAPIUserOTPLogin(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) + w = app.doJSON(http.MethodPost, "/api/v1/users", map[string]any{"username": "lisi", "email": "ls@example.com"}, ck) + if w.Code != http.StatusOK { + t.Fatalf("create user status = %d, body=%s", w.Code, w.Body.String()) + } + + // 生成图形验证码(直接经 store,模拟用户看到验证码) + cap, err := app.captchas.New() + if err != nil { + t.Fatalf("captcha new: %v", err) + } + + // 发送 OTP + w = app.doJSON(http.MethodPost, "/api/v1/auth/otp/send", map[string]any{ + "username": "ext_lisi", "captcha_id": cap.ID, "captcha_code": cap.Text, + }) + if w.Code != http.StatusOK { + t.Fatalf("otp send status = %d, body=%s", w.Code, w.Body.String()) + } + + // CLI 通道取同一验证码 + code, err := app.otps.Current(context.Background(), "ext_lisi") + if err != nil { + t.Fatalf("otp current: %v", err) + } + + // OTP 登录 + w = app.doJSON(http.MethodPost, "/api/v1/auth/otp/login", map[string]string{"username": "ext_lisi", "code": code}) + if w.Code != http.StatusOK { + t.Fatalf("otp login status = %d, body=%s", w.Code, w.Body.String()) + } + userCk := sessionCookie(t, w) + + // 外部用户 me + w = app.doJSON(http.MethodGet, "/api/v1/auth/me", nil, userCk) + if w.Code != http.StatusOK { + t.Fatalf("user me status = %d", w.Code) + } + if m := decodeBody(t, w); m["data"].(map[string]any)["user_type"] != "user" { + t.Fatalf("user me body = %s", w.Body.String()) + } + + // 外部用户访问 admin 路由应 403 + w = app.doJSON(http.MethodGet, "/api/v1/users", nil, userCk) + if w.Code != http.StatusForbidden { + t.Fatalf("user access admin status = %d, want 403", w.Code) + } + + // 登出 + w = app.doJSON(http.MethodPost, "/api/v1/auth/logout", nil, userCk) + if w.Code != http.StatusOK { + t.Fatalf("logout status = %d", w.Code) + } + w = app.doJSON(http.MethodGet, "/api/v1/auth/me", nil, userCk) + if w.Code != http.StatusUnauthorized { + t.Fatalf("me after logout status = %d, want 401", w.Code) + } +} + +func TestAPIAdminForgotReset(t *testing.T) { + app := setupTestApp(t) + w := app.doJSON(http.MethodPost, "/api/v1/auth/admin/forgot", map[string]string{"username": "root"}) + if w.Code != http.StatusOK { + t.Fatalf("forgot status = %d", w.Code) + } + if app.mailer.lastTo != "root@example.com" { + t.Fatalf("reset mail to = %q", app.mailer.lastTo) + } + idx := bytes.Index([]byte(app.mailer.lastBody), []byte("token=")) + if idx < 0 { + t.Fatalf("reset link missing token: %s", app.mailer.lastBody) + } + token := app.mailer.lastBody[idx+6:] + token = token[:bytes.IndexByte([]byte(token), '\n')] + w = app.doJSON(http.MethodPost, "/api/v1/auth/admin/reset", map[string]string{"token": token, "new_password": "NewPassw0rd"}) + if w.Code != http.StatusOK { + t.Fatalf("reset status = %d, body=%s", w.Code, w.Body.String()) + } + // 新密码可登录 + w = app.doJSON(http.MethodPost, "/api/v1/auth/admin/login", map[string]string{"username": "root", "password": "NewPassw0rd"}) + if w.Code != http.StatusOK { + t.Fatalf("login with new password status = %d", w.Code) + } +} + +func itoa(u uint) string { + return strconv.FormatUint(uint64(u), 10) +} diff --git a/internal/api/auth.go b/internal/api/auth.go new file mode 100644 index 0000000..255c025 --- /dev/null +++ b/internal/api/auth.go @@ -0,0 +1,195 @@ +package api + +import ( + "encoding/base64" + "errors" + "net/http" + + "github.com/gin-gonic/gin" + + "ws_usernode/internal/auth" + "ws_usernode/internal/config" + "ws_usernode/internal/service" +) + +// AuthHandler 认证接口:图形验证码、OTP 双通道登录、管理员登录、密码重置、会话。 +type AuthHandler struct { + svc *service.AuthService + cfg *config.Config +} + +// Captcha GET /auth/captcha —— 获取图形验证码(id + base64 PNG)。 +func (h *AuthHandler) Captcha(c *gin.Context) { + id, png, err := h.svc.NewCaptcha() + if err != nil { + fail(c, http.StatusInternalServerError, "验证码生成失败") + return + } + ok(c, gin.H{ + "captcha_id": id, + "image": "data:image/png;base64," + base64.StdEncoding.EncodeToString(png), + }) +} + +// OTPSendRequest 外部用户请求 OTP。 +type OTPSendRequest struct { + Username string `json:"username" binding:"required"` // 含或不含 ext_ 前缀 + CaptchaID string `json:"captcha_id" binding:"required"` + CaptchaCode string `json:"captcha_code" binding:"required"` +} + +// OTPSend POST /auth/otp/send —— 图形验证码前置,生成 OTP 并发邮件(失败不阻断)。 +func (h *AuthHandler) OTPSend(c *gin.Context) { + var req OTPSendRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + err := h.svc.UserOTPSend(c.Request.Context(), req.Username, req.CaptchaID, req.CaptchaCode, c.ClientIP()) + if err != nil { + switch { + case errors.Is(err, service.ErrCaptchaFailed): + fail(c, http.StatusBadRequest, "图形验证码错误") + case errors.Is(err, service.ErrUserUnavailable): + fail(c, http.StatusNotFound, "用户不存在或不可用") + case errors.Is(err, auth.ErrCooldown): + fail(c, http.StatusTooManyRequests, "发送冷却中,请稍后重试") + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + ok(c, gin.H{"status": "sent"}) +} + +// OTPLoginRequest 外部用户 OTP 登录。 +type OTPLoginRequest struct { + Username string `json:"username" binding:"required"` + Code string `json:"code" binding:"required"` // 6 位 OTP +} + +// OTPLogin POST /auth/otp/login —— OTP 校验并建立 cookie 会话。 +func (h *AuthHandler) OTPLogin(c *gin.Context) { + var req OTPLoginRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + sid, err := h.svc.UserOTPLogin(c.Request.Context(), req.Username, req.Code, c.ClientIP(), c.Request.UserAgent()) + if err != nil { + switch { + case errors.Is(err, service.ErrUserUnavailable): + fail(c, http.StatusNotFound, "用户不存在或不可用") + case errors.Is(err, auth.ErrInvalidCode): + fail(c, http.StatusUnauthorized, "验证码错误或已过期") + case errors.Is(err, auth.ErrTooManyFails): + fail(c, http.StatusTooManyRequests, "失败次数过多,请稍后再试") + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + setSessionCookie(c, sid, h.cfg.Server.SessionTTL, h.cfg.App.Env == "production") + ok(c, gin.H{"session": "created"}) +} + +// AdminLoginRequest 管理员登录。 +type AdminLoginRequest struct { + Username string `json:"username" binding:"required"` + Password string `json:"password" binding:"required"` +} + +// AdminLogin POST /auth/admin/login —— 管理员用户名+口令登录(含失败限速)。 +func (h *AuthHandler) AdminLogin(c *gin.Context) { + var req AdminLoginRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + sid, err := h.svc.AdminLogin(c.Request.Context(), req.Username, req.Password, c.ClientIP(), c.Request.UserAgent()) + if err != nil { + switch { + case errors.Is(err, service.ErrRateLimited): + fail(c, http.StatusTooManyRequests, "尝试次数过多,请稍后再试") + case errors.Is(err, service.ErrBadCredentials): + fail(c, http.StatusUnauthorized, "用户名或密码错误") + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + setSessionCookie(c, sid, h.cfg.Server.SessionTTL, h.cfg.App.Env == "production") + ok(c, gin.H{"session": "created"}) +} + +// AdminForgotRequest 管理员忘记密码。 +type AdminForgotRequest struct { + Username string `json:"username" binding:"required"` +} + +// AdminForgot POST /auth/admin/forgot —— 发送密码重置邮件(用户不存在也返回成功,防枚举)。 +func (h *AuthHandler) AdminForgot(c *gin.Context) { + var req AdminForgotRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + if err := h.svc.AdminForgot(c.Request.Context(), req.Username, c.ClientIP()); err != nil { + fail(c, http.StatusInternalServerError, err.Error()) + return + } + ok(c, gin.H{"status": "sent"}) +} + +// AdminResetRequest 通过令牌重置密码。 +type AdminResetRequest struct { + Token string `json:"token" binding:"required"` + NewPassword string `json:"new_password" binding:"required"` +} + +// AdminReset POST /auth/admin/reset —— 校验令牌并重置密码。 +func (h *AuthHandler) AdminReset(c *gin.Context) { + var req AdminResetRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + err := h.svc.AdminReset(c.Request.Context(), req.Token, req.NewPassword, c.ClientIP()) + if err != nil { + switch { + case errors.Is(err, auth.ErrResetTokenInvalid): + fail(c, http.StatusBadRequest, "重置令牌无效或已过期") + case errors.Is(err, service.ErrWeakPassword): + fail(c, http.StatusBadRequest, err.Error()) + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + ok(c, gin.H{"status": "reset"}) +} + +// Logout POST /auth/logout —— 登出(会话删除 + cookie 清除)。 +func (h *AuthHandler) Logout(c *gin.Context) { + sess := sessionFrom(c) + if sess != nil { + _ = h.svc.Logout(c.Request.Context(), sess.ID) + } + setSessionCookie(c, "", 0, h.cfg.App.Env == "production") + ok(c, gin.H{"status": "logged_out"}) +} + +// Me GET /auth/me —— 当前会话主体信息。 +func (h *AuthHandler) Me(c *gin.Context) { + sess := sessionFrom(c) + if sess == nil { + fail(c, http.StatusUnauthorized, "未登录") + return + } + info, err := h.svc.Me(c.Request.Context(), sess.ID) + if err != nil { + fail(c, http.StatusUnauthorized, "会话失效或已过期") + return + } + ok(c, info) +} diff --git a/internal/api/cookie.go b/internal/api/cookie.go new file mode 100644 index 0000000..7b16dff --- /dev/null +++ b/internal/api/cookie.go @@ -0,0 +1,40 @@ +package api + +import ( + "net/http" + "time" + + "github.com/gin-gonic/gin" +) + +// SessionCookieName 会话 cookie 名称。 +const SessionCookieName = "usernode_session" + +// SessionContextKey 会话在 gin context 中的键(router 中间件注入)。 +const SessionContextKey = "auth_session" + +// setSessionCookie 写入会话 cookie(HttpOnly/SameSite=Lax;maxAge<=0 时清除)。 +func setSessionCookie(c *gin.Context, sid string, ttl time.Duration, secure bool) { + maxAge := int(ttl.Seconds()) + if sid == "" { + maxAge = -1 + } + http.SetCookie(c.Writer, &http.Cookie{ + Name: SessionCookieName, + Value: sid, + Path: "/", + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + MaxAge: maxAge, + Secure: secure, + }) +} + +// SessionIDFromCookie 从请求 cookie 读取会话 ID。 +func SessionIDFromCookie(c *gin.Context) string { + v, err := c.Cookie(SessionCookieName) + if err != nil { + return "" + } + return v +} diff --git a/internal/api/handler.go b/internal/api/handler.go index 416d710..3dc7bca 100644 --- a/internal/api/handler.go +++ b/internal/api/handler.go @@ -1,36 +1,57 @@ // Package api 为 HTTP handler 层(RESTful v1)。 -// M0 提供健康检查与模块路由骨架;各模块 handler 在对应里程碑填充。 +// M1 覆盖认证(管理员/外部用户登录、会话、OTP)与用户管理 CRUD。 package api import ( "net/http" "runtime" + "strconv" "time" "github.com/gin-gonic/gin" + "ws_usernode/internal/auth" + "ws_usernode/internal/config" "ws_usernode/internal/service" ) // Handler 聚合各模块 handler,作为路由注册的挂载点。 type Handler struct { Health *HealthHandler - Admin *AdminHandler + Auth *AuthHandler User *UserHandler - // Auth / Keys / Approval / Audit / Settings 等模块在 M1~M4 填充 + + authSvc *service.AuthService + auditSvc *service.AuditService } -// New 创建 handler 集合。M0 阶段部分服务可为 nil,路由只挂已实现模块。 -func New(adminSvc *service.AdminService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler { +// New 创建 handler 集合。 +func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler { h := &Handler{ - Health: &HealthHandler{startedAt: time.Now()}, - Admin: &AdminHandler{svc: adminSvc}, - User: &UserHandler{svc: userSvc}, + Health: &HealthHandler{startedAt: time.Now()}, + authSvc: authSvc, + auditSvc: auditSvc, } - _ = auditSvc + h.Auth = &AuthHandler{svc: authSvc, cfg: cfg} + h.User = &UserHandler{svc: userSvc, cfg: cfg, h: h} return h } +// audit 记录管理操作审计(append-only)。actor 来自会话中间件。 +func (h *Handler) audit(c *gin.Context, action, resourceType, resourceID string, detail any, result string) { + var actorID uint + var actorName string + if sess := sessionFrom(c); sess != nil { + actorID = sess.RefID + if info, err := h.authSvc.Me(c.Request.Context(), sess.ID); err == nil { + actorName = info.Username + } else { + actorName = sess.UserType + "#" + strconv.FormatUint(uint64(sess.RefID), 10) + } + } + _ = h.auditSvc.Record(c.Request.Context(), actorID, actorName, action, resourceType, resourceID, detail, c.ClientIP(), result) +} + // HealthHandler 健康检查。 type HealthHandler struct { startedAt time.Time @@ -39,13 +60,23 @@ type HealthHandler struct { func (h *HealthHandler) Healthz(c *gin.Context) { c.JSON(http.StatusOK, gin.H{ "status": "ok", - "version": "0.1.0-m0", + "version": "0.2.0-m1", "uptime": time.Since(h.startedAt).String(), "go": runtime.Version(), "timestamp": time.Now().UTC().Format(time.RFC3339), }) } +// sessionFrom 返回会话中间件注入的会话(未登录时为 nil)。 +func sessionFrom(c *gin.Context) *auth.Session { + if v, ok := c.Get(SessionContextKey); ok { + if s, ok := v.(*auth.Session); ok { + return s + } + } + return nil +} + // ok 统一成功响应。 func ok(c *gin.Context, data any) { c.JSON(http.StatusOK, gin.H{"data": data}) diff --git a/internal/api/user.go b/internal/api/user.go index 7cfe2d5..5a72628 100644 --- a/internal/api/user.go +++ b/internal/api/user.go @@ -1,17 +1,24 @@ package api import ( + "context" + "errors" "net/http" - "strings" + "strconv" + "time" "github.com/gin-gonic/gin" + "ws_usernode/internal/config" + "ws_usernode/internal/model" "ws_usernode/internal/service" ) -// UserHandler 外部用户接口(列表/详情/创建等,M1 填充 CRUD 与系统操作)。 +// UserHandler 外部用户接口(列表/详情/创建/更新/禁用/启用/延期/删除,admin)。 type UserHandler struct { svc *service.UserService + cfg *config.Config + h *Handler // 访问审计 helper } // UserCreateRequest 管理员创建外部用户请求。 @@ -25,31 +32,187 @@ type UserCreateRequest struct { // Create 管理员创建外部用户(自动建系统账号)。 func (h *UserHandler) Create(c *gin.Context) { - if h.svc == nil { - fail(c, http.StatusNotImplemented, "用户服务未初始化") - return - } var req UserCreateRequest if err := c.ShouldBindJSON(&req); err != nil { fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) return } - u, err := h.svc.Create(c.Request.Context(), req.Username, req.Email, req.Supervisor, req.Purpose, req.TTLDays*86400) + createdBy := uint(0) + if sess := sessionFrom(c); sess != nil { + createdBy = sess.RefID + } + u, err := h.svc.Create(c.Request.Context(), req.Username, req.Email, req.Supervisor, req.Purpose, + time.Duration(req.TTLDays)*24*time.Hour, createdBy) if err != nil { + h.h.audit(c, "user.create", "user", "", map[string]any{"username": req.Username, "err": err.Error()}, model.ResultFailed) switch { - case strings.Contains(err.Error(), "用户名"): + case errors.Is(err, service.ErrUserExists): + fail(c, http.StatusConflict, err.Error()) + default: fail(c, http.StatusBadRequest, err.Error()) - case strings.Contains(err.Error(), "已存在"): + } + return + } + h.h.audit(c, "user.create", "user", strconv.FormatUint(uint64(u.ID), 10), map[string]any{"username": u.Username}, model.ResultSuccess) + ok(c, gin.H{"id": u.ID, "username": u.Username, "status": u.Status, "expire_at": u.ExpireAt}) +} + +// List 用户列表(分页/筛选:status、supervisor)。 +func (h *UserHandler) List(c *gin.Context) { + page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) + pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) + users, total, err := h.svc.List(c.Request.Context(), service.UserFilter{ + Status: c.Query("status"), + Supervisor: c.Query("supervisor"), + Page: page, + PageSize: pageSize, + }) + if err != nil { + fail(c, http.StatusInternalServerError, err.Error()) + return + } + ok(c, gin.H{"total": total, "items": users}) +} + +// Get 用户详情。 +func (h *UserHandler) Get(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的用户 ID") + return + } + u, err := h.svc.GetByID(c.Request.Context(), uint(id)) + if err != nil { + fail(c, http.StatusNotFound, err.Error()) + return + } + ok(c, u) +} + +// UserUpdateRequest 更新外部用户信息(仅更新提供的字段;邮箱仅管理员可改)。 +type UserUpdateRequest struct { + Email *string `json:"email"` + Supervisor *string `json:"supervisor"` + Purpose *string `json:"purpose"` +} + +// Update PATCH /users/:id。 +func (h *UserHandler) Update(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的用户 ID") + return + } + var req UserUpdateRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + u, err := h.svc.Update(c.Request.Context(), uint(id), req.Email, req.Supervisor, req.Purpose) + if err != nil { + h.h.audit(c, "user.update", "user", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed) + if errors.Is(err, service.ErrUserNotFound) { + fail(c, http.StatusNotFound, err.Error()) + return + } + fail(c, http.StatusBadRequest, err.Error()) + return + } + h.h.audit(c, "user.update", "user", c.Param("id"), map[string]any{"email": req.Email, "supervisor": req.Supervisor, "purpose": req.Purpose}, model.ResultSuccess) + ok(c, u) +} + +// Disable POST /users/:id/disable —— 禁用(清空 authorized_keys,SSH 立即失效)。 +func (h *UserHandler) Disable(c *gin.Context) { + h.setStatus(c, "user.disable", model.UserStatusDisabled, h.svc.Disable) +} + +// Enable POST /users/:id/enable —— 启用(按 DB 密钥状态恢复)。 +func (h *UserHandler) Enable(c *gin.Context) { + h.setStatus(c, "user.enable", model.UserStatusActive, h.svc.Enable) +} + +// setStatus 复用禁用/启用的公共流程(解析 ID、调 service、审计)。 +func (h *UserHandler) setStatus(c *gin.Context, action, wantStatus string, fn func(ctx context.Context, id uint) error) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的用户 ID") + return + } + if err := fn(c, uint(id)); err != nil { + h.h.audit(c, action, "user", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed) + switch { + case errors.Is(err, service.ErrUserNotFound): + fail(c, http.StatusNotFound, err.Error()) + case errors.Is(err, service.ErrUserExpired): + fail(c, http.StatusConflict, err.Error()) + case errors.Is(err, service.ErrSystemAccountMissing): fail(c, http.StatusConflict, err.Error()) default: fail(c, http.StatusInternalServerError, err.Error()) } return } - ok(c, gin.H{"id": u.ID, "username": u.Username, "status": u.Status}) + h.h.audit(c, action, "user", c.Param("id"), map[string]any{"status": wantStatus}, model.ResultSuccess) + ok(c, gin.H{"status": wantStatus}) } -// List 用户列表(M1 实现分页筛选)。 -func (h *UserHandler) List(c *gin.Context) { - fail(c, http.StatusNotImplemented, "用户列表将在 M1 实现") +// ExtendRequest 延期请求。 +type ExtendRequest struct { + Days int `json:"days"` // 0 表示用配置默认(90 天) +} + +// Extend POST /users/:id/extend —— 延长有效期;已过期用户在回收期内可恢复。 +func (h *UserHandler) Extend(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的用户 ID") + return + } + var req ExtendRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + u, err := h.svc.GetByID(c.Request.Context(), uint(id)) + if err != nil { + fail(c, http.StatusNotFound, err.Error()) + return + } + if err := h.svc.Extend(c.Request.Context(), uint(id), req.Days); 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 +} + +// Delete DELETE /users/:id —— 删除并回收(系统账号 + 家目录 + 密钥,保留审计)。 +func (h *UserHandler) Delete(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的用户 ID") + return + } + if err := h.svc.Delete(c.Request.Context(), uint(id)); err != nil { + h.h.audit(c, "user.delete", "user", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed) + switch { + case errors.Is(err, service.ErrUserNotFound): + fail(c, http.StatusNotFound, err.Error()) + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + h.h.audit(c, "user.delete", "user", c.Param("id"), nil, model.ResultSuccess) + ok(c, gin.H{"status": "deleted"}) } diff --git a/internal/auth/captcha.go b/internal/auth/captcha.go index 5793182..8ba6c6b 100644 --- a/internal/auth/captcha.go +++ b/internal/auth/captcha.go @@ -14,10 +14,11 @@ var ErrCaptchaInvalid = errors.New("auth: 图形验证码错误") // Captcha 图形验证码(防机器人,登录前置)。 type Captcha struct { ID string - Text string // M1 生成图像渲染,此处仅存文本 + Text string } -// CaptchaStore 为图形验证码存储(M1 实现图像渲染)。 +// CaptchaStore 为图形验证码存储。单实例内存实现为默认; +// 多实例部署需改 DB/Redis(PLAN §6 注明)。 type CaptchaStore interface { // New 生成一个验证码并返回其 ID。 New() (*Captcha, error) @@ -27,6 +28,7 @@ type CaptchaStore interface { // MemoryCaptchaStore 单实例内存实现。 type MemoryCaptchaStore struct { + ttl time.Duration mu sync.Mutex entries map[string]*captchaEntry } @@ -37,8 +39,8 @@ type captchaEntry struct { } // NewMemoryCaptchaStore 创建内存图形验证码存储。 -func NewMemoryCaptchaStore() *MemoryCaptchaStore { - return &MemoryCaptchaStore{entries: make(map[string]*captchaEntry)} +func NewMemoryCaptchaStore(ttl time.Duration) *MemoryCaptchaStore { + return &MemoryCaptchaStore{ttl: ttl, entries: make(map[string]*captchaEntry)} } func (s *MemoryCaptchaStore) New() (*Captcha, error) { @@ -52,7 +54,7 @@ func (s *MemoryCaptchaStore) New() (*Captcha, error) { } s.mu.Lock() defer s.mu.Unlock() - s.entries[id] = &captchaEntry{text: text, expiresAt: time.Now().Add(5 * time.Minute)} + s.entries[id] = &captchaEntry{text: text, expiresAt: time.Now().Add(s.ttl)} return &Captcha{ID: id, Text: text}, nil } diff --git a/internal/auth/captcha_img.go b/internal/auth/captcha_img.go new file mode 100644 index 0000000..587d4ca --- /dev/null +++ b/internal/auth/captcha_img.go @@ -0,0 +1,131 @@ +package auth + +import ( + "bytes" + "fmt" + "image" + "image/color" + "image/draw" + "image/png" + "math/rand" + "time" +) + +// 验证码图像渲染:内置 5x7 点阵数字 + 噪点 + 干扰线,不依赖第三方字体库 +// (PLAN §4:内置生成,不依赖第三方服务)。 + +// digitGlyphs 为 0-9 的 5x7 点阵,每行低 5 位表示该行像素。 +var digitGlyphs = [10][7]byte{ + {0b01110, 0b10001, 0b10011, 0b10101, 0b11001, 0b10001, 0b01110}, // 0 + {0b00100, 0b01100, 0b00100, 0b00100, 0b00100, 0b00100, 0b01110}, // 1 + {0b01110, 0b10001, 0b00001, 0b00010, 0b00100, 0b01000, 0b11111}, // 2 + {0b11111, 0b00010, 0b00100, 0b00010, 0b00001, 0b10001, 0b01110}, // 3 + {0b00010, 0b00110, 0b01010, 0b10010, 0b11111, 0b00010, 0b00010}, // 4 + {0b11111, 0b10000, 0b11110, 0b00001, 0b00001, 0b10001, 0b01110}, // 5 + {0b00110, 0b01000, 0b10000, 0b11110, 0b10001, 0b10001, 0b01110}, // 6 + {0b11111, 0b00001, 0b00010, 0b00100, 0b01000, 0b01000, 0b01000}, // 7 + {0b01110, 0b10001, 0b10001, 0b01110, 0b10001, 0b10001, 0b01110}, // 8 + {0b01110, 0b10001, 0b10001, 0b01111, 0b00001, 0b00010, 0b01100}, // 9 +} + +const ( + glyphW = 5 + glyphH = 7 + scale = 3 // 点阵放大倍数 + charGap = 4 // 字符间距(像素) + edgePad = 6 // 画布边距 +) + +// RenderCaptchaPNG 将 4 位数字验证码渲染为 PNG 字节流。 +// 仅接受数字字符,其余返回错误。 +func RenderCaptchaPNG(text string) ([]byte, error) { + if len(text) == 0 || len(text) > 8 { + return nil, fmt.Errorf("auth: captcha text length must be 1~8") + } + for _, r := range text { + if r < '0' || r > '9' { + return nil, fmt.Errorf("auth: captcha text must be digits") + } + } + canvasW := edgePad*2 + len(text)*glyphW*scale + (len(text)-1)*charGap + canvasH := edgePad*2 + glyphH*scale + img := image.NewRGBA(image.Rect(0, 0, canvasW, canvasH)) + draw.Draw(img, img.Bounds(), &image.Uniform{C: color.RGBA{248, 250, 252, 255}}, image.Point{}, draw.Src) + + rng := rand.New(rand.NewSource(time.Now().UnixNano())) + + // 噪点(浅灰,稀疏) + for i := 0; i < canvasW*canvasH/7; i++ { + x, y := rng.Intn(canvasW), rng.Intn(canvasH) + g := uint8(150 + rng.Intn(90)) + img.Set(x, y, color.RGBA{g, g, g, 255}) + } + + // 干扰线(穿过字符区域的浅色斜线) + for i := 0; i < 3; i++ { + g := uint8(160 + rng.Intn(80)) + c := color.RGBA{g, g, g, 255} + x1, y1 := rng.Intn(canvasW/2), rng.Intn(canvasH) + x2, y2 := canvasW/2+rng.Intn(canvasW/2), rng.Intn(canvasH) + drawLine(img, x1, y1, x2, y2, c) + } + + // 逐字符绘制(颜色随机取深色系) + inkPalette := []color.RGBA{ + {40, 60, 110, 255}, {120, 45, 45, 255}, {30, 90, 60, 255}, {80, 60, 110, 255}, + } + for i, r := range text { + glyph := digitGlyphs[r-'0'] + ink := inkPalette[rng.Intn(len(inkPalette))] + x0 := edgePad + i*(glyphW*scale+charGap) + y0 := edgePad + for row := 0; row < glyphH; row++ { + for col := 0; col < glyphW; col++ { + if glyph[row]&(1<<(glyphW-1-col)) != 0 { + fillRect(img, x0+col*scale, y0+row*scale, scale, scale, ink) + } + } + } + } + + var buf bytes.Buffer + if err := png.Encode(&buf, img); err != nil { + return nil, fmt.Errorf("auth: encode captcha png: %w", err) + } + return buf.Bytes(), nil +} + +// fillRect 填充实心矩形。 +func fillRect(img *image.RGBA, x, y, w, h int, c color.RGBA) { + for dy := 0; dy < h; dy++ { + for dx := 0; dx < w; dx++ { + img.Set(x+dx, y+dy, c) + } + } +} + +// drawLine 使用 DDA 算法画线。 +func drawLine(img *image.RGBA, x1, y1, x2, y2 int, c color.RGBA) { + steps := abs(x2-x1) + if d := abs(y2 - y1); d > steps { + steps = d + } + if steps == 0 { + img.Set(x1, y1, c) + return + } + for i := 0; i <= steps; i++ { + x := x1 + (x2-x1)*i/steps + y := y1 + (y2-y1)*i/steps + if x >= 0 && x < img.Bounds().Dx() && y >= 0 && y < img.Bounds().Dy() { + img.Set(x, y, c) + } + } +} + +func abs(n int) int { + if n < 0 { + return -n + } + return n +} diff --git a/internal/auth/captcha_test.go b/internal/auth/captcha_test.go new file mode 100644 index 0000000..d8a70e3 --- /dev/null +++ b/internal/auth/captcha_test.go @@ -0,0 +1,115 @@ +package auth + +import ( + "context" + "image/png" + "strings" + "testing" + "time" +) + +func TestRenderCaptchaPNG(t *testing.T) { + pngBytes, err := RenderCaptchaPNG("4837") + if err != nil { + t.Fatalf("render: %v", err) + } + img, err := png.Decode(strings.NewReader(string(pngBytes))) + if err != nil { + t.Fatalf("decode png: %v", err) + } + if img.Bounds().Dx() < 50 || img.Bounds().Dy() < 20 { + t.Fatalf("canvas too small: %dx%d", img.Bounds().Dx(), img.Bounds().Dy()) + } +} + +func TestRenderCaptchaPNGInvalid(t *testing.T) { + if _, err := RenderCaptchaPNG("12a4"); err == nil { + t.Fatal("expected error for non-digit text") + } + if _, err := RenderCaptchaPNG(""); err == nil { + t.Fatal("expected error for empty text") + } +} + +func TestMemoryCaptchaStore(t *testing.T) { + s := NewMemoryCaptchaStore(5 * time.Minute) + cap, err := s.New() + if err != nil { + t.Fatalf("new: %v", err) + } + if len(cap.Text) != 4 { + t.Fatalf("captcha text len = %d, want 4", len(cap.Text)) + } + if !s.Verify(cap.ID, cap.Text) { + t.Fatal("verify should succeed") + } + if s.Verify(cap.ID, cap.Text) { + t.Fatal("captcha must be one-time") + } +} + +func TestMemoryCaptchaStoreExpired(t *testing.T) { + s := NewMemoryCaptchaStore(-time.Second) // 立即过期 + cap, err := s.New() + if err != nil { + t.Fatalf("new: %v", err) + } + if s.Verify(cap.ID, cap.Text) { + t.Fatal("expired captcha should fail") + } +} + +func TestRateLimiter(t *testing.T) { + l := NewRateLimiter(3, time.Minute) + key := "admin-login:admin" + if !l.Allow(key) { + t.Fatal("first attempt should be allowed") + } + for i := 0; i < 3; i++ { + l.RecordFailure(key) + } + if l.Allow(key) { + t.Fatal("should be locked after max failures") + } + l.Reset(key) + if !l.Allow(key) { + t.Fatal("should be allowed after reset") + } +} + +func TestDBResetTokenStore(t *testing.T) { + db := testDB(t) + s := NewDBResetTokenStore(db) + ctx := context.Background() + + token, err := s.Create(ctx, 42, time.Minute, "127.0.0.1") + if err != nil { + t.Fatalf("create: %v", err) + } + if token == "" { + t.Fatal("token should not be empty") + } + adminID, err := s.Consume(ctx, token) + if err != nil { + t.Fatalf("consume: %v", err) + } + if adminID != 42 { + t.Fatalf("adminID = %d, want 42", adminID) + } + // 一次性 + if _, err := s.Consume(ctx, token); err != ErrResetTokenInvalid { + t.Fatalf("second consume err = %v, want ErrResetTokenInvalid", err) + } + // 无效令牌 + if _, err := s.Consume(ctx, "not-a-token"); err != ErrResetTokenInvalid { + t.Fatalf("bad token err = %v, want ErrResetTokenInvalid", err) + } + // 过期令牌 + dbToken, err := s.Create(ctx, 7, -time.Minute, "") + if err != nil { + t.Fatalf("create expired: %v", err) + } + if _, err := s.Consume(ctx, dbToken); err != ErrResetTokenInvalid { + t.Fatalf("expired token err = %v, want ErrResetTokenInvalid", err) + } +} diff --git a/internal/auth/otp.go b/internal/auth/otp.go index 445a7e1..4b05ab5 100644 --- a/internal/auth/otp.go +++ b/internal/auth/otp.go @@ -1,15 +1,22 @@ // Package auth 提供认证相关能力:OTP(双通道)、会话、bcrypt、图形验证码。 // // OTP 双通道对齐:邮件发送与 CLI 获取共用同一 OTPStore(同一验证码、同一 -// 10 分钟有效期、同一 60s 冷却与失败限速),邮件失败不阻断 CLI 通道。 -// 单实例用内存存储;多实例需改为 DB/Redis(PLAN §6 注明)。 +// 有效期、同一冷却与失败限速),邮件失败不阻断 CLI 通道。邮件通道经 +// Send 生成验证码,CLI 通道优先经 Current 复用同一验证码,无有效码时才 +// 触发 Send(仍受同一冷却约束)。验证码落 DB(otp_codes 表),单实例部署 +// 即可保证跨进程(HTTP 服务与 CLI 子命令)共享;多实例需改 DB 行锁/Redis +// (PLAN §6)。 package auth import ( + "context" "errors" "sync" "time" + "gorm.io/gorm" + + "ws_usernode/internal/model" "ws_usernode/internal/pkg" ) @@ -21,20 +28,141 @@ var ( ) const ( - otpCodeLen = 6 - maxFailures = 5 // 单账号连续失败限速阈值 - failureWin = 10 * time.Minute // 失败计数窗口 + otpCodeLen = 6 ) -// OTPStore 为 OTP 验证码存储。内存实现为单实例默认实现。 +// OTP 失败限速默认参数(单账号连续失败阈值与计数窗口)。 +const ( + DefaultMaxFailures = 5 + DefaultFailureWin = 10 * time.Minute +) + +// OTPStore 为 OTP 验证码存储。DB 实现(DBOTPStore)为生产默认, +// MemoryOTPStore 供测试与单进程内嵌场景使用。 type OTPStore interface { - // Send 为 username 生成新验证码(覆盖旧码)。冷却期内调用返回 ErrCooldown。 - // 邮件与 CLI 双通道都走该方法,保证对齐。 - Send(username string, ttl, cooldown time.Duration) (string, error) + // Send 为 username 生成新验证码并覆盖旧码(邮件通道)。冷却期内返回 ErrCooldown。 + Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) + // Current 返回当前有效(未过期、未消费)验证码,供 CLI 通道复用同一验证码。 + // 无有效验证码返回 ErrInvalidCode。 + Current(ctx context.Context, username string) (string, error) // Verify 校验验证码并一次性消费。失败累计计数(达到阈值返回 ErrTooManyFails)。 - Verify(username, code string) (bool, error) + Verify(ctx context.Context, username, code string) (bool, error) // Failures 返回 username 当前失败计数。 - Failures(username string) (int, error) + Failures(ctx context.Context, username string) (int, error) +} + +// DBOTPStore 基于 model.OTPCode 的存储实现,每用户一行(username 唯一)。 +type DBOTPStore struct { + db *gorm.DB + maxFailures int + failureWin time.Duration +} + +// NewDBOTPStore 创建 DB OTP 存储。 +func NewDBOTPStore(db *gorm.DB, maxFailures int, failureWin time.Duration) *DBOTPStore { + return &DBOTPStore{db: db, maxFailures: maxFailures, failureWin: failureWin} +} + +func (s *DBOTPStore) Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) { + now := time.Now() + var row model.OTPCode + err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error + // 注意:First 的 err 不能复用给后续语句,避免被覆盖导致走错分支 + exists := err == nil + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return "", err + } + if exists && now.Before(row.CooldownUntil) { + return "", ErrCooldown + } + code, randErr := pkg.RandomDigits(otpCodeLen) + if randErr != nil { + return "", randErr + } + updates := map[string]any{ + "code": code, + "expires_at": now.Add(ttl), + "cooldown_until": now.Add(cooldown), + "failures": 0, + "failed_at": nil, + "consumed_at": nil, + } + if exists { + err = s.db.WithContext(ctx).Model(&row).Updates(updates).Error + } else { + err = s.db.WithContext(ctx).Create(&model.OTPCode{ + Username: username, + Code: code, + ExpiresAt: now.Add(ttl), + CooldownUntil: now.Add(cooldown), + }).Error + } + if err != nil { + return "", err + } + return code, nil +} + +func (s *DBOTPStore) Current(ctx context.Context, username string) (string, error) { + var row model.OTPCode + if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return "", ErrInvalidCode + } + return "", err + } + if row.ConsumedAt != nil || time.Now().After(row.ExpiresAt) { + return "", ErrInvalidCode + } + return row.Code, nil +} + +func (s *DBOTPStore) Verify(ctx context.Context, username, code string) (bool, error) { + now := time.Now() + var row model.OTPCode + if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, ErrInvalidCode + } + return false, err + } + if row.ConsumedAt != nil || now.After(row.ExpiresAt) { + return false, ErrInvalidCode + } + // 失败计数窗口:超出窗口则重置计数 + if row.FailedAt != nil && now.Sub(*row.FailedAt) > s.failureWin { + row.Failures = 0 + row.FailedAt = nil + } + if row.Failures >= s.maxFailures { + return false, ErrTooManyFails + } + if row.Code != code { + row.Failures++ + f := now + row.FailedAt = &f + _ = s.db.WithContext(ctx).Model(&row).Updates(map[string]any{"failures": row.Failures, "failed_at": row.FailedAt}).Error + if row.Failures >= s.maxFailures { + return false, ErrTooManyFails + } + return false, ErrInvalidCode + } + consumed := now + if err := s.db.WithContext(ctx).Model(&row).Update("consumed_at", &consumed).Error; err != nil { + return false, err + } + return true, nil +} + +func (s *DBOTPStore) Failures(ctx context.Context, username string) (int, error) { + var row model.OTPCode + if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, nil + } + return 0, err + } + return row.Failures, nil } type otpEntry struct { @@ -44,7 +172,7 @@ type otpEntry struct { failures int } -// MemoryOTPStore 为单实例内存实现。 +// MemoryOTPStore 为单进程内存实现(测试/内嵌场景)。 type MemoryOTPStore struct { mu sync.Mutex entries map[string]*otpEntry @@ -55,9 +183,10 @@ func NewMemoryOTPStore() *MemoryOTPStore { return &MemoryOTPStore{entries: make(map[string]*otpEntry)} } -func (s *MemoryOTPStore) Send(username string, ttl, cooldown time.Duration) (string, error) { +func (s *MemoryOTPStore) Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) { s.mu.Lock() defer s.mu.Unlock() + _ = ctx now := time.Now() if e, ok := s.entries[username]; ok && now.Before(e.cooldownAt) { @@ -76,9 +205,21 @@ func (s *MemoryOTPStore) Send(username string, ttl, cooldown time.Duration) (str return code, nil } -func (s *MemoryOTPStore) Verify(username, code string) (bool, error) { +func (s *MemoryOTPStore) Current(ctx context.Context, username string) (string, error) { s.mu.Lock() defer s.mu.Unlock() + _ = ctx + e, ok := s.entries[username] + if !ok || time.Now().After(e.expiresAt) { + return "", ErrInvalidCode + } + return e.code, nil +} + +func (s *MemoryOTPStore) Verify(ctx context.Context, username, code string) (bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + _ = ctx e, ok := s.entries[username] if !ok { @@ -89,12 +230,12 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) { delete(s.entries, username) return false, ErrInvalidCode } - if e.failures >= maxFailures { + if e.failures >= DefaultMaxFailures { return false, ErrTooManyFails } if e.code != code { e.failures++ - if e.failures >= maxFailures { + if e.failures >= DefaultMaxFailures { return false, ErrTooManyFails } return false, ErrInvalidCode @@ -103,9 +244,10 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) { return true, nil } -func (s *MemoryOTPStore) Failures(username string) (int, error) { +func (s *MemoryOTPStore) Failures(ctx context.Context, username string) (int, error) { s.mu.Lock() defer s.mu.Unlock() + _ = ctx if e, ok := s.entries[username]; ok { return e.failures, nil } diff --git a/internal/auth/otp_test.go b/internal/auth/otp_test.go new file mode 100644 index 0000000..7b19e4c --- /dev/null +++ b/internal/auth/otp_test.go @@ -0,0 +1,119 @@ +package auth + +import ( + "context" + "testing" + "time" + + "gorm.io/gorm" + + "ws_usernode/internal/model" +) + +func testDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := model.Open("sqlite", ":memory:", false) + if err != nil { + t.Fatalf("open test db: %v", err) + } + if err := model.Migrate(db); err != nil { + t.Fatalf("migrate: %v", err) + } + return db +} + +func TestDBOTPStoreSendCurrentVerify(t *testing.T) { + db := testDB(t) + s := NewDBOTPStore(db, DefaultMaxFailures, DefaultFailureWin) + ctx := context.Background() + const user = "ext_zhangsan" + + code, err := s.Send(ctx, user, 10*time.Minute, time.Minute) + if err != nil { + t.Fatalf("send: %v", err) + } + if len(code) != 6 { + t.Fatalf("code length = %d, want 6", len(code)) + } + // 双通道对齐:Current 复用同一验证码 + cur, err := s.Current(ctx, user) + if err != nil { + t.Fatalf("current: %v", err) + } + if cur != code { + t.Fatalf("current = %q, want %q (双通道必须同一验证码)", cur, code) + } + // 校验成功 + ok, err := s.Verify(ctx, user, code) + if err != nil || !ok { + t.Fatalf("verify = %v/%v, want true/nil", ok, err) + } + // 一次性:再次校验失败 + ok, err = s.Verify(ctx, user, code) + if err != ErrInvalidCode { + t.Fatalf("second verify err = %v, want ErrInvalidCode", err) + } + if ok { + t.Fatal("second verify should fail") + } + // Current 也应失败(已消费) + if _, err := s.Current(ctx, user); err != ErrInvalidCode { + t.Fatalf("current after consume err = %v, want ErrInvalidCode", err) + } +} + +func TestDBOTPStoreCooldown(t *testing.T) { + db := testDB(t) + s := NewDBOTPStore(db, DefaultMaxFailures, DefaultFailureWin) + ctx := context.Background() + if _, err := s.Send(ctx, "ext_lisi", time.Minute, time.Minute); err != nil { + t.Fatalf("send: %v", err) + } + if _, err := s.Send(ctx, "ext_lisi", time.Minute, time.Minute); err != ErrCooldown { + t.Fatalf("second send err = %v, want ErrCooldown", err) + } +} + +func TestDBOTPStoreFailuresAndWindow(t *testing.T) { + db := testDB(t) + s := NewDBOTPStore(db, 3, 10*time.Minute) + ctx := context.Background() + code, err := s.Send(ctx, "ext_wangwu", time.Minute, 0) + if err != nil { + t.Fatalf("send: %v", err) + } + // 3 次错误后进入限速 + for i := 0; i < 3; i++ { + if _, err := s.Verify(ctx, "ext_wangwu", "000000"); err != ErrInvalidCode && err != ErrTooManyFails { + t.Fatalf("verify(%d) err = %v", i, err) + } + } + if _, err := s.Verify(ctx, "ext_wangwu", code); err != ErrTooManyFails { + t.Fatalf("verify after max failures err = %v, want ErrTooManyFails", err) + } + f, err := s.Failures(ctx, "ext_wangwu") + if err != nil { + t.Fatalf("failures: %v", err) + } + if f != 3 { + t.Fatalf("failures = %d, want 3", f) + } +} + +func TestMemoryOTPStore(t *testing.T) { + s := NewMemoryOTPStore() + ctx := context.Background() + code, err := s.Send(ctx, "ext_test", time.Minute, time.Second) + if err != nil { + t.Fatalf("send: %v", err) + } + if cur, err := s.Current(ctx, "ext_test"); err != nil || cur != code { + t.Fatalf("current = %q/%v, want %q/nil", cur, err, code) + } + if _, err := s.Send(ctx, "ext_test", time.Minute, time.Second); err != ErrCooldown { + t.Fatalf("send during cooldown err = %v, want ErrCooldown", err) + } + if ok, err := s.Verify(ctx, "ext_test", code); err != nil || !ok { + t.Fatalf("verify = %v/%v, want true/nil", ok, err) + } +} diff --git a/internal/auth/ratelimit.go b/internal/auth/ratelimit.go new file mode 100644 index 0000000..38b0535 --- /dev/null +++ b/internal/auth/ratelimit.go @@ -0,0 +1,68 @@ +package auth + +import ( + "sync" + "time" +) + +// RateLimiter 内存登录限速器:按 key(如用户名)统计连续失败次数, +// 达到阈值后锁定 lockFor 时长,锁定期内 Allow 返回 false。 +// 单实例内存实现即可满足(登录失败限速无跨进程一致性要求)。 +type RateLimiter struct { + mu sync.Mutex + max int + lockFor time.Duration + entries map[string]*rlEntry + now func() time.Time // 可注入时钟(测试) +} + +type rlEntry struct { + failures int + lockedUntil time.Time +} + +// NewRateLimiter 创建限速器。 +func NewRateLimiter(max int, lockFor time.Duration) *RateLimiter { + return &RateLimiter{max: max, lockFor: lockFor, entries: make(map[string]*rlEntry), now: time.Now} +} + +// Allow 返回 key 当前是否允许继续尝试。 +func (l *RateLimiter) Allow(key string) bool { + l.mu.Lock() + defer l.mu.Unlock() + e, ok := l.entries[key] + if !ok { + return true + } + if now := l.now(); now.Before(e.lockedUntil) { + return false + } else if e.failures >= l.max { + // 锁定已过期:重置计数,允许重试 + delete(l.entries, key) + return true + } + return true +} + +// RecordFailure 记录一次失败;达到阈值后进入锁定。 +func (l *RateLimiter) RecordFailure(key string) { + l.mu.Lock() + defer l.mu.Unlock() + now := l.now() + e, ok := l.entries[key] + if !ok || now.After(e.lockedUntil) && e.failures >= l.max { + e = &rlEntry{} + l.entries[key] = e + } + e.failures++ + if e.failures >= l.max { + e.lockedUntil = now.Add(l.lockFor) + } +} + +// Reset 清除 key 的失败记录(登录成功后调用)。 +func (l *RateLimiter) Reset(key string) { + l.mu.Lock() + defer l.mu.Unlock() + delete(l.entries, key) +} diff --git a/internal/auth/reset.go b/internal/auth/reset.go new file mode 100644 index 0000000..cde51a5 --- /dev/null +++ b/internal/auth/reset.go @@ -0,0 +1,76 @@ +package auth + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "time" + + "gorm.io/gorm" + + "ws_usernode/internal/model" + "ws_usernode/internal/pkg" +) + +// ErrResetTokenInvalid 表示重置令牌无效、已使用或已过期。 +var ErrResetTokenInvalid = errors.New("auth: 重置令牌无效或已过期") + +// ResetTokenStore 为管理员密码重置令牌存储(邮件重置)。 +type ResetTokenStore interface { + // Create 生成令牌并存储其哈希,返回令牌明文(仅经邮件/日志发出)。 + Create(ctx context.Context, adminID uint, ttl time.Duration, ip string) (string, error) + // Consume 校验令牌并标记已使用,返回对应的管理员 ID。 + Consume(ctx context.Context, token string) (uint, error) +} + +// DBResetTokenStore 基于 model.PasswordResetToken 的存储实现。 +type DBResetTokenStore struct { + db *gorm.DB +} + +// NewDBResetTokenStore 创建重置令牌存储。 +func NewDBResetTokenStore(db *gorm.DB) *DBResetTokenStore { + return &DBResetTokenStore{db: db} +} + +func (s *DBResetTokenStore) Create(ctx context.Context, adminID uint, ttl time.Duration, ip string) (string, error) { + token, err := pkg.RandomHex(24) + if err != nil { + return "", err + } + row := model.PasswordResetToken{ + AdminID: adminID, + TokenHash: hashToken(token), + ExpiresAt: time.Now().Add(ttl), + IP: ip, + } + if err := s.db.WithContext(ctx).Create(&row).Error; err != nil { + return "", err + } + return token, nil +} + +func (s *DBResetTokenStore) Consume(ctx context.Context, token string) (uint, error) { + var row model.PasswordResetToken + if err := s.db.WithContext(ctx).Where("token_hash = ?", hashToken(token)).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, ErrResetTokenInvalid + } + return 0, err + } + if row.UsedAt != nil || time.Now().After(row.ExpiresAt) { + return 0, ErrResetTokenInvalid + } + now := time.Now() + if err := s.db.WithContext(ctx).Model(&row).Update("used_at", &now).Error; err != nil { + return 0, err + } + return row.AdminID, nil +} + +// hashToken 计算令牌的 SHA-256 摘要(令牌本身为高熵随机串,无需加盐)。 +func hashToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} diff --git a/internal/config/config.go b/internal/config/config.go index ba9a96e..47edf85 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -28,13 +28,15 @@ type Config struct { Database DatabaseConfig `toml:"database"` Log LogConfig `toml:"log"` Policy PolicyConfig `toml:"policy"` + Auth AuthConfig `toml:"auth"` SMTP SMTPConfig `toml:"smtp"` System SystemConfig `toml:"system"` } type AppConfig struct { - Name string `toml:"name"` - Env string `toml:"env"` // development / production + Name string `toml:"name"` + Env string `toml:"env"` // development / production + BaseURL string `toml:"base_url"` // 对外访问地址(邮件中的链接使用) } type ServerConfig struct { @@ -62,6 +64,13 @@ type PolicyConfig struct { OTPCooldown time.Duration `toml:"otp_cooldown"` // OTP 发送冷却 } +// AuthConfig 认证与防暴力参数。 +type AuthConfig struct { + MaxLoginFailures int `toml:"max_login_failures"` // 管理员登录连续失败阈值,达到后锁定 + LockDuration time.Duration `toml:"lock_duration"` // 失败达到阈值后的锁定时长 + CaptchaTTL time.Duration `toml:"captcha_ttl"` // 图形验证码有效期 +} + type SMTPConfig struct { Host string `toml:"host"` Port int `toml:"port"` @@ -72,21 +81,22 @@ type SMTPConfig struct { // SystemConfig 为系统账号操作层的本地实现配置(sudoers 白名单模式)。 type SystemConfig struct { - Sudo bool `toml:"sudo"` // 是否通过 sudo -n 执行系统命令;开发环境 false 时 dry-run - UserPrefix string `toml:"user_prefix"` // 外部用户系统账号前缀,默认 ext_ - Group string `toml:"group"` // 外部用户所属组,默认 external - Shell string `toml:"shell"` // 默认 shell - HomeBase string `toml:"home_base"` // 家目录基路径 - AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖) + Sudo bool `toml:"sudo"` // 是否通过 sudo -n 执行系统命令(生产) + DryRun bool `toml:"dry_run"` // true = 只打印计划命令不执行(开发演练);false 且 sudo=false 时直接执行(容器/测试用户验证) + UserPrefix string `toml:"user_prefix"` // 外部用户系统账号前缀,默认 ext_ + Group string `toml:"group"` // 外部用户所属组,默认 external + Shell string `toml:"shell"` // 默认 shell + HomeBase string `toml:"home_base"` // 家目录基路径 + AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖) } // Default 返回带开发环境默认值的配置,作为 config.example.toml 与未配置项的兜底。 func Default() *Config { return &Config{ - App: AppConfig{Name: "ws_usernode", Env: "development"}, + App: AppConfig{Name: "ws_usernode", Env: "development", BaseURL: "http://127.0.0.1:8080"}, Server: ServerConfig{ - Listen: "127.0.0.1:8080", - SessionTTL: 24 * time.Hour, + Listen: "127.0.0.1:8080", + SessionTTL: 24 * time.Hour, TrustedProxies: []string{"127.0.0.1", "::1"}, }, Database: DatabaseConfig{Driver: "sqlite", DSN: "data/usernode.db"}, @@ -98,9 +108,17 @@ func Default() *Config { OTPTTL: 10 * time.Minute, OTPCooldown: 60 * time.Second, }, + Auth: AuthConfig{ + MaxLoginFailures: 5, + LockDuration: 15 * time.Minute, + CaptchaTTL: 5 * time.Minute, + }, SMTP: SMTPConfig{Port: 587}, System: SystemConfig{ + // 开发默认 dry-run:未配置 config 直接跑 serve 时只打印计划,避免误操作系统账号。 + // 生产必须显式 dry_run=false 且 sudo=true(见 deploy/sudoers.example)。 Sudo: false, + DryRun: true, UserPrefix: "ext_", Group: "external", Shell: "/bin/sh", diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4b36ab3..bf72e98 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -31,6 +31,12 @@ func TestLoadDefault(t *testing.T) { if cfg.System.UserPrefix != "ext_" { t.Errorf("default user_prefix = %q", cfg.System.UserPrefix) } + if cfg.Auth.MaxLoginFailures != 5 || cfg.Auth.LockDuration != 15*time.Minute { + t.Errorf("default auth = %+v", cfg.Auth) + } + if !cfg.System.DryRun { + t.Error("default dry_run should be true (安全默认)") + } } func TestLoadFileOverrides(t *testing.T) { @@ -62,7 +68,11 @@ func TestEnvOverrides(t *testing.T) { t.Setenv("USERNODE_DATABASE_DSN", "u:p@tcp(h:3306)/db") t.Setenv("USERNODE_POLICY_OTPTTL", "5m") t.Setenv("USERNODE_SYSTEM_SUDO", "true") + t.Setenv("USERNODE_SYSTEM_DRY_RUN", "false") t.Setenv("USERNODE_SERVER_TRUSTED_PROXIES", "10.0.0.1, 10.0.0.2") + t.Setenv("USERNODE_AUTH_MAX_LOGIN_FAILURES", "3") + t.Setenv("USERNODE_AUTH_LOCK_DURATION", "5m") + t.Setenv("USERNODE_APP_BASE_URL", "https://un.example.com") cfg, err := LoadDefault() if err != nil { @@ -77,6 +87,12 @@ func TestEnvOverrides(t *testing.T) { if !cfg.System.Sudo { t.Error("system.sudo should be true") } + if cfg.Auth.MaxLoginFailures != 3 || cfg.Auth.LockDuration != 5*time.Minute { + t.Errorf("auth = %+v", cfg.Auth) + } + if cfg.App.BaseURL != "https://un.example.com" { + t.Errorf("base_url = %q", cfg.App.BaseURL) + } if len(cfg.Server.TrustedProxies) != 2 || cfg.Server.TrustedProxies[0] != "10.0.0.1" { t.Errorf("trusted_proxies = %v", cfg.Server.TrustedProxies) } diff --git a/internal/mail/mailer.go b/internal/mail/mailer.go new file mode 100644 index 0000000..fd32c8e --- /dev/null +++ b/internal/mail/mailer.go @@ -0,0 +1,118 @@ +// Package mail 提供邮件发送抽象。M1 提供基础直发(net/smtp + STARTTLS), +// 发送队列/重试/失败记录(mail_logs)在 M3 完善;SMTP 未配置时退化为 +// LogMailer(仅打印,不阻断业务——OTP 邮件失败不阻断 CLI 通道)。 +package mail + +import ( + "context" + "crypto/tls" + "encoding/base64" + "fmt" + "log/slog" + "net" + "net/smtp" + "strconv" + "strings" + + "ws_usernode/internal/config" +) + +// Mailer 邮件发送接口。 +type Mailer interface { + // Send 发送一封纯文本邮件;失败返回错误,由调用方决定是否阻断。 + Send(ctx context.Context, to, subject, body string) error +} + +// LogMailer 在 SMTP 未配置时替代实现:把邮件内容打到日志。 +// 生产部署必须配置 SMTP,此时日志不落邮件正文。 +type LogMailer struct { + log *slog.Logger +} + +// NewLogMailer 创建日志邮件实现。 +func NewLogMailer(log *slog.Logger) *LogMailer { + return &LogMailer{log: log} +} + +func (m *LogMailer) Send(_ context.Context, to, subject, body string) error { + m.log.Warn("mail: smtp 未配置,邮件内容仅写入日志", + "to", to, "subject", subject, "body", body) + return nil +} + +// SMTPMailer 基于 net/smtp 的基础直发(STARTTLS + 可选 AUTH LOGIN/PLAIN)。 +type SMTPMailer struct { + cfg config.SMTPConfig +} + +// NewSMTPMailer 创建 SMTP 邮件实现。 +func NewSMTPMailer(cfg config.SMTPConfig) *SMTPMailer { + return &SMTPMailer{cfg: cfg} +} + +// New 按配置选择实现:SMTP host 为空时返回 LogMailer。 +func New(cfg config.SMTPConfig, log *slog.Logger) Mailer { + if cfg.Host == "" { + return NewLogMailer(log) + } + return NewSMTPMailer(cfg) +} + +func (m *SMTPMailer) Send(ctx context.Context, to, subject, body string) error { + addr := net.JoinHostPort(m.cfg.Host, strconv.Itoa(m.cfg.Port)) + + var d net.Dialer + conn, err := d.DialContext(ctx, "tcp", addr) + if err != nil { + return fmt.Errorf("mail: dial %s: %w", addr, err) + } + defer conn.Close() + + c, err := smtp.NewClient(conn, m.cfg.Host) + if err != nil { + return fmt.Errorf("mail: smtp client: %w", err) + } + defer c.Close() + + if err := c.StartTLS(&tls.Config{ServerName: m.cfg.Host}); err != nil { + return fmt.Errorf("mail: starttls: %w", err) + } + if m.cfg.Username != "" { + if err := c.Auth(smtp.PlainAuth("", m.cfg.Username, m.cfg.Password, m.cfg.Host)); err != nil { + return fmt.Errorf("mail: auth: %w", err) + } + } + if err := c.Mail(m.cfg.From); err != nil { + return fmt.Errorf("mail: mail from: %w", err) + } + if err := c.Rcpt(to); err != nil { + return fmt.Errorf("mail: rcpt: %w", err) + } + w, err := c.Data() + if err != nil { + return fmt.Errorf("mail: data: %w", err) + } + msg := buildMessage(m.cfg.From, to, subject, body) + if _, err := w.Write([]byte(msg)); err != nil { + w.Close() + return fmt.Errorf("mail: write body: %w", err) + } + if err := w.Close(); err != nil { + return fmt.Errorf("mail: close body: %w", err) + } + return c.Quit() +} + +// buildMessage 构造 RFC 5322 消息体(UTF-8 主题 base64 编码,正文 UTF-8)。 +func buildMessage(from, to, subject, body string) string { + var b strings.Builder + b.WriteString("From: " + from + "\r\n") + b.WriteString("To: " + to + "\r\n") + b.WriteString("Subject: =?UTF-8?B?" + base64.StdEncoding.EncodeToString([]byte(subject)) + "?=\r\n") + b.WriteString("MIME-Version: 1.0\r\n") + b.WriteString("Content-Type: text/plain; charset=UTF-8\r\n") + b.WriteString("Content-Transfer-Encoding: 8bit\r\n") + b.WriteString("\r\n") + b.WriteString(body + "\r\n") + return b.String() +} diff --git a/internal/model/model.go b/internal/model/model.go index ad105ac..250ebe5 100644 --- a/internal/model/model.go +++ b/internal/model/model.go @@ -141,11 +141,43 @@ type MailLog struct { UpdatedAt time.Time `json:"updated_at"` } +// OTPCode 外部用户 OTP 验证码(DB 存储,邮件与 CLI 双通道共享)。 +// +// 每用户一行(username 唯一):邮件通道经 Send 生成,CLI 通道经 Current +// 读取同一验证码,保证"同一验证码、同一有效期、同一冷却与失败限速"(PLAN §2.2)。 +// 验证码为 6 位数字,生命周期短(10 分钟)且一次性消费,存储明文以便 CLI +// 复用返回;并发写由单实例部署的串行事务保证(多实例需改 DB 行锁/Redis,PLAN §6)。 +type OTPCode struct { + ID uint `gorm:"primaryKey" json:"id"` + Username string `gorm:"size:64;uniqueIndex;not null" json:"username"` + Code string `gorm:"size:16;not null" json:"-"` + ExpiresAt time.Time `gorm:"index;not null" json:"expires_at"` + CooldownUntil time.Time `json:"cooldown_until"` // Send 冷却截止 + Failures int `json:"failures"` + FailedAt *time.Time `json:"failed_at"` // 失败计数窗口起点(窗口内达阈值限速) + ConsumedAt *time.Time `json:"consumed_at"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// PasswordResetToken 管理员密码重置令牌(邮件重置)。DB 只存哈希, +// 明文令牌仅经邮件/日志发送给管理员,一次有效。 +type PasswordResetToken struct { + ID uint `gorm:"primaryKey" json:"id"` + AdminID uint `gorm:"index;not null" json:"admin_id"` + TokenHash string `gorm:"size:64;not null" json:"-"` + ExpiresAt time.Time `gorm:"index;not null" json:"expires_at"` + UsedAt *time.Time `json:"used_at"` + IP string `gorm:"size:64" json:"ip"` + CreatedAt time.Time `json:"created_at"` +} + // AllModels 供 AutoMigrate 使用的全部模型。 func AllModels() []any { return []any{ &AdminUser{}, &User{}, &SSHKey{}, &Approval{}, &AuditLog{}, &Session{}, &Setting{}, &MailLog{}, + &OTPCode{}, &PasswordResetToken{}, } } diff --git a/internal/router/middleware.go b/internal/router/middleware.go new file mode 100644 index 0000000..f8da1c4 --- /dev/null +++ b/internal/router/middleware.go @@ -0,0 +1,46 @@ +package router + +import ( + "net/http" + + "github.com/gin-gonic/gin" + + "ws_usernode/internal/api" + "ws_usernode/internal/auth" +) + +// sessionMiddleware 从 cookie 解析会话并注入 gin context;未登录返回 401。 +// 会话为 DB 存储,任何节点实例均可校验(兼容多实例)。 +func sessionMiddleware(sessions auth.SessionStore) gin.HandlerFunc { + return func(c *gin.Context) { + sid := api.SessionIDFromCookie(c) + if sid == "" { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "未登录"}) + return + } + sess, err := sessions.Get(c.Request.Context(), sid) + if err != nil { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "会话失效或已过期"}) + return + } + c.Set(api.SessionContextKey, sess) + c.Next() + } +} + +// requireUserType 校验会话主体类型(admin / user)。空串表示任意登录主体。 +func requireUserType(userType string) gin.HandlerFunc { + return func(c *gin.Context) { + sess, ok := c.Get(api.SessionContextKey) + if !ok { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "未登录"}) + return + } + s := sess.(*auth.Session) + if userType != "" && s.UserType != userType { + c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "权限不足"}) + return + } + c.Next() + } +} diff --git a/internal/router/router.go b/internal/router/router.go index 1148c55..5547021 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -9,13 +9,14 @@ import ( "github.com/gin-gonic/gin" "ws_usernode/internal/api" + "ws_usernode/internal/auth" "ws_usernode/internal/config" "ws_usernode/internal/webui" ) // New 构建根 router:API v1 + 前端静态资源(go:embed)。 // production 模式启用 gin.ReleaseMode;否则启用调试模式与开发日志。 -func New(cfg *config.Config, h *api.Handler, log *slog.Logger) *gin.Engine { +func New(cfg *config.Config, h *api.Handler, sessions auth.SessionStore, log *slog.Logger) *gin.Engine { if cfg.App.Env == "production" { gin.SetMode(gin.ReleaseMode) } @@ -28,21 +29,34 @@ func New(cfg *config.Config, h *api.Handler, log *slog.Logger) *gin.Engine { // RESTful API v1 v1 := r.Group("/api/v1") { - auth := v1.Group("/auth") + authGrp := v1.Group("/auth") { - // M1:captcha / otp/send / otp/login / admin/login / logout / me - auth.GET("/captcha", notImplemented("图形验证码(M1)")) - auth.POST("/otp/send", notImplemented("OTP 发送(M1)")) - auth.POST("/otp/login", notImplemented("OTP 登录(M1)")) - auth.POST("/admin/login", notImplemented("管理员登录(M1)")) - auth.POST("/logout", notImplemented("登出(M1)")) - auth.GET("/me", notImplemented("当前会话(M1)")) + authGrp.GET("/captcha", h.Auth.Captcha) + authGrp.POST("/otp/send", h.Auth.OTPSend) + authGrp.POST("/otp/login", h.Auth.OTPLogin) + authGrp.POST("/admin/login", h.Auth.AdminLogin) + authGrp.POST("/admin/forgot", h.Auth.AdminForgot) + authGrp.POST("/admin/reset", h.Auth.AdminReset) + // 需要会话(管理员或外部用户) + authed := authGrp.Group("", sessionMiddleware(sessions)) + authed.POST("/logout", h.Auth.Logout) + authed.GET("/me", h.Auth.Me) } - v1.GET("/users", h.User.List) - v1.POST("/users", h.User.Create) - v1.POST("/users/:id/disable", notImplemented("禁用用户(M1)")) - v1.POST("/users/:id/enable", notImplemented("启用用户(M1)")) - v1.POST("/users/:id/extend", notImplemented("延期(M1)")) + + // 用户管理(admin) + users := v1.Group("/users", sessionMiddleware(sessions), requireUserType(auth.SessionUserAdmin)) + { + users.GET("", h.User.List) + users.POST("", h.User.Create) + users.GET("/:id", h.User.Get) + users.PATCH("/:id", h.User.Update) + users.POST("/:id/disable", h.User.Disable) + users.POST("/:id/enable", h.User.Enable) + users.POST("/:id/extend", h.User.Extend) + users.DELETE("/:id", h.User.Delete) + } + + // 后续里程碑 v1.POST("/approvals", notImplemented("提交申请(M3)")) v1.GET("/approvals", notImplemented("申请列表(M3)")) v1.POST("/approvals/:id/review", notImplemented("审批(M3)")) diff --git a/internal/service/admin.go b/internal/service/admin.go index d20aed8..5d65275 100644 --- a/internal/service/admin.go +++ b/internal/service/admin.go @@ -15,9 +15,10 @@ import ( // 管理员服务错误。 var ( - ErrAdminExists = errors.New("service: 管理员已存在") - ErrAdminNotFound = errors.New("service: 管理员不存在") - ErrWeakPassword = errors.New("service: 密码过弱(至少 8 位,需含字母与数字)") + ErrAdminExists = errors.New("service: 管理员已存在") + ErrAdminNotFound = errors.New("service: 管理员不存在") + ErrWeakPassword = errors.New("service: 密码过弱(至少 8 位,需含字母与数字)") + ErrBadCredentials = errors.New("service: 用户名或密码错误") ) // AdminService 管理端账号服务(CLI admin create / reset-password 与登录共用)。 @@ -94,6 +95,43 @@ func (s *AdminService) GetByUsername(ctx context.Context, username string) (*mod return &adm, nil } +// GetByID 按 ID 查询管理员。 +func (s *AdminService) GetByID(ctx context.Context, id uint) (*model.AdminUser, error) { + var adm model.AdminUser + if err := s.db.WithContext(ctx).First(&adm, "id = ?", id).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrAdminNotFound + } + return nil, err + } + return &adm, nil +} + +// GetByEmail 按邮箱查询管理员(用于密码重置邮件)。 +func (s *AdminService) GetByEmail(ctx context.Context, email string) (*model.AdminUser, error) { + var adm model.AdminUser + if err := s.db.WithContext(ctx).First(&adm, "email = ?", email).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrAdminNotFound + } + return nil, err + } + return &adm, nil +} + +// Login 校验管理员用户名与口令(bcrypt)。失败统一返回 ErrBadCredentials, +// 不暴露账号是否存在。登录限速由 AuthService 的 RateLimiter 处理。 +func (s *AdminService) Login(ctx context.Context, username, password string) (*model.AdminUser, error) { + adm, err := s.GetByUsername(ctx, strings.TrimSpace(username)) + if err != nil { + return nil, ErrBadCredentials + } + if !auth.VerifyPassword(adm.PasswordHash, password) { + return nil, ErrBadCredentials + } + return adm, nil +} + // validatePassword 校验管理员密码强度(骨架阶段基础规则,M1 可加策略)。 func validatePassword(p string) error { if len(p) < 8 { diff --git a/internal/service/auth.go b/internal/service/auth.go new file mode 100644 index 0000000..cc238ba --- /dev/null +++ b/internal/service/auth.go @@ -0,0 +1,235 @@ +package service + +import ( + "context" + "errors" + "log/slog" + "strings" + "time" + + "gorm.io/gorm" + + "ws_usernode/internal/auth" + "ws_usernode/internal/config" + "ws_usernode/internal/mail" + "ws_usernode/internal/model" +) + +// 认证服务错误。 +var ( + ErrRateLimited = errors.New("service: 尝试次数过多,请稍后再试") + ErrCaptchaFailed = errors.New("service: 图形验证码错误") + ErrUserUnavailable = errors.New("service: 用户不存在或不可用") +) + +// AuthService 认证服务:管理员/外部用户登录、会话、图形验证码、OTP 双通道、 +// 密码重置。持有各存储与邮件发送器,作为 api 层与底层存储之间的桥。 +type AuthService struct { + db *gorm.DB + cfg *config.Config + otps auth.OTPStore + captchas auth.CaptchaStore + sessions auth.SessionStore + resets auth.ResetTokenStore + limiter *auth.RateLimiter + mailer mail.Mailer + users *UserService + admins *AdminService + audit *AuditService + log *slog.Logger +} + +// NewAuthService 组装认证服务。 +func NewAuthService(db *gorm.DB, cfg *config.Config, otps auth.OTPStore, captchas auth.CaptchaStore, + sessions auth.SessionStore, resets auth.ResetTokenStore, limiter *auth.RateLimiter, + mailer mail.Mailer, users *UserService, admins *AdminService, audit *AuditService, log *slog.Logger) *AuthService { + return &AuthService{ + db: db, cfg: cfg, otps: otps, captchas: captchas, sessions: sessions, + resets: resets, limiter: limiter, mailer: mailer, users: users, + admins: admins, audit: audit, log: log, + } +} + +// NewCaptcha 生成图形验证码并渲染 PNG 图像。 +func (s *AuthService) NewCaptcha() (id string, png []byte, err error) { + cap, err := s.captchas.New() + if err != nil { + return "", nil, err + } + png, err = auth.RenderCaptchaPNG(cap.Text) + if err != nil { + return "", nil, err + } + return cap.ID, png, nil +} + +// createSession 建立 cookie 会话(DB 存储)。 +func (s *AuthService) createSession(ctx context.Context, userType string, refID uint, ip, userAgent string) (string, error) { + return s.sessions.Create(ctx, userType, refID, s.cfg.Server.SessionTTL, ip, userAgent) +} + +// AdminLogin 管理员用户名+口令登录,返回会话 ID。 +// 连续失败达到阈值(config auth.max_login_failures)后锁定 lock_duration。 +func (s *AuthService) AdminLogin(ctx context.Context, username, password, ip, userAgent string) (string, error) { + key := "admin-login:" + strings.TrimSpace(username) + if !s.limiter.Allow(key) { + _ = s.audit.Record(ctx, 0, username, "admin.login", "admin", "", map[string]any{"locked": true}, ip, model.ResultFailed) + return "", ErrRateLimited + } + adm, err := s.admins.Login(ctx, username, password) + if err != nil { + s.limiter.RecordFailure(key) + _ = s.audit.Record(ctx, 0, username, "admin.login", "admin", "", nil, ip, model.ResultFailed) + return "", ErrBadCredentials + } + s.limiter.Reset(key) + sid, err := s.createSession(ctx, auth.SessionUserAdmin, adm.ID, ip, userAgent) + if err != nil { + return "", err + } + _ = s.audit.Record(ctx, adm.ID, adm.Username, "admin.login", "admin", "", nil, ip, model.ResultSuccess) + return sid, nil +} + +// Logout 注销会话(管理员或外部用户通用)。 +func (s *AuthService) Logout(ctx context.Context, sessionID string) error { + return s.sessions.Delete(ctx, sessionID) +} + +// AdminForgot 发送密码重置邮件(无 SMTP 时退化为日志输出)。 +// 用户不存在时也返回成功,避免账号枚举。 +func (s *AuthService) AdminForgot(ctx context.Context, username, ip string) error { + username = strings.TrimSpace(username) + adm, err := s.admins.GetByUsername(ctx, username) + if err != nil { + s.log.Warn("admin forgot: user not found (not revealed)", "username", username) + return nil + } + token, err := s.resets.Create(ctx, adm.ID, 30*time.Minute, ip) + if err != nil { + return err + } + link := strings.TrimRight(s.cfg.App.BaseURL, "/") + "/reset?token=" + token + body := "您正在重置管理员密码。请在 30 分钟内打开以下链接完成重置:\n\n" + link + + "\n\n如非本人操作请忽略本邮件。也可由管理员通过 CLI `usernode admin reset-password` 重置。" + if err := s.mailer.Send(ctx, adm.Email, "重置密码", body); err != nil { + s.log.Warn("admin forgot: mail send failed", "err", err) + } + _ = s.audit.Record(ctx, adm.ID, adm.Username, "admin.forgot", "admin", "", nil, ip, model.ResultSuccess) + return nil +} + +// AdminReset 通过令牌重置管理员密码,并使该管理员既有会话全部失效。 +func (s *AuthService) AdminReset(ctx context.Context, token, newPassword string, ip string) error { + adminID, err := s.resets.Consume(ctx, token) + if err != nil { + return err + } + adm, err := s.admins.GetByID(ctx, adminID) + if err != nil { + return err + } + if err := validatePassword(newPassword); err != nil { + return err + } + hash, err := auth.HashPassword(newPassword) + if err != nil { + return err + } + if err := s.db.WithContext(ctx).Model(&model.AdminUser{}).Where("id = ?", adm.ID).Update("password_hash", hash).Error; err != nil { + return err + } + // 重置后吊销该管理员全部会话 + _ = s.db.WithContext(ctx).Where("user_type = ? AND ref_id = ?", auth.SessionUserAdmin, adm.ID).Delete(&model.Session{}).Error + _ = s.audit.Record(ctx, adm.ID, adm.Username, "admin.reset", "admin", "", nil, ip, model.ResultSuccess) + return nil +} + +// UserOTPSend 外部用户申请 OTP:图形验证码前置,生成验证码并邮件发送。 +// 邮件失败不阻断(CLI 通道兜底)。冷却/失败限速与 CLI 通道共享同一存储。 +func (s *AuthService) UserOTPSend(ctx context.Context, username, captchaID, captchaAnswer, ip string) error { + if !s.captchas.Verify(captchaID, captchaAnswer) { + return ErrCaptchaFailed + } + u, err := s.users.GetByUsername(ctx, username) + if err != nil { + return ErrUserUnavailable + } + if u.Status != model.UserStatusActive { + return ErrUserUnavailable + } + code, err := s.otps.Send(ctx, u.Username, s.cfg.Policy.OTPTTL, s.cfg.Policy.OTPCooldown) + if err != nil { + return err + } + body := "您的登录验证码是:" + code + + "\n有效期 " + s.cfg.Policy.OTPTTL.String() + ",请勿向他人泄露。\n" + + "如未收到邮件,可通过 CLI 子命令 `usernode user otp --username " + u.Username + "` 获取同一验证码。" + if err := s.mailer.Send(ctx, u.Email, "登录验证码", body); err != nil { + // OTP 邮件失败不阻断登录(PLAN F5);CLI 通道仍可获取同一验证码 + s.log.Warn("otp mail send failed, cli channel remains available", "username", u.Username, "err", err) + } + _ = s.audit.Record(ctx, u.ID, u.Username, "user.otp.send", "user", "", nil, ip, model.ResultSuccess) + return nil +} + +// UserOTPLogin 外部用户 OTP 登录,返回会话 ID。 +func (s *AuthService) UserOTPLogin(ctx context.Context, username, code, ip, userAgent string) (string, error) { + u, err := s.users.GetByUsername(ctx, username) + if err != nil { + return "", ErrUserUnavailable + } + if u.Status != model.UserStatusActive { + return "", ErrUserUnavailable + } + ok, err := s.otps.Verify(ctx, u.Username, code) + if err != nil { + _ = s.audit.Record(ctx, u.ID, u.Username, "user.otp.login", "user", "", map[string]any{"err": err.Error()}, ip, model.ResultFailed) + return "", err + } + if !ok { + return "", auth.ErrInvalidCode + } + now := time.Now() + _ = s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", u.ID).Update("last_login_at", &now).Error + sid, err := s.createSession(ctx, auth.SessionUserUser, u.ID, ip, userAgent) + if err != nil { + return "", err + } + _ = s.audit.Record(ctx, u.ID, u.Username, "user.otp.login", "user", "", nil, ip, model.ResultSuccess) + return sid, nil +} + +// MeInfo 当前会话对应的主体信息。 +type MeInfo struct { + UserType string `json:"user_type"` + ID uint `json:"id"` + Username string `json:"username"` + Email string `json:"email"` +} + +// Me 根据会话 ID 返回当前登录主体。 +func (s *AuthService) Me(ctx context.Context, sessionID string) (*MeInfo, error) { + sess, err := s.sessions.Get(ctx, sessionID) + if err != nil { + return nil, err + } + switch sess.UserType { + case auth.SessionUserAdmin: + adm, err := s.admins.GetByID(ctx, sess.RefID) + if err != nil { + return nil, err + } + return &MeInfo{UserType: sess.UserType, ID: adm.ID, Username: adm.Username, Email: adm.Email}, nil + case auth.SessionUserUser: + u, err := s.users.GetByID(ctx, sess.RefID) + if err != nil { + return nil, err + } + return &MeInfo{UserType: sess.UserType, ID: u.ID, Username: u.Username, Email: u.Email}, nil + } + return nil, ErrSessionInvalid +} + +// ErrSessionInvalid 表示会话类型未知。 +var ErrSessionInvalid = errors.New("service: 会话无效") diff --git a/internal/service/auth_test.go b/internal/service/auth_test.go new file mode 100644 index 0000000..b25c8bd --- /dev/null +++ b/internal/service/auth_test.go @@ -0,0 +1,230 @@ +package service + +import ( + "context" + "errors" + "io" + "log/slog" + "strings" + "testing" + + "ws_usernode/internal/auth" +) + +// testLogger 丢弃日志输出的测试 logger。 +func testLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +// recordingMailer 捕获最近一封邮件,便于从重置邮件中提取令牌。 +type recordingMailer struct { + lastTo string + lastSubject string + lastBody string +} + +func (m *recordingMailer) Send(_ context.Context, to, subject, body string) error { + m.lastTo = to + m.lastSubject = subject + m.lastBody = body + return nil +} + +// newTestAuthService 组装一套完整的认证服务(DB + 内存验证码 + DB OTP/会话/令牌)。 +func newTestAuthService(t *testing.T) (*AuthService, *fakeSys, *recordingMailer) { + t.Helper() + db := testDB(t) + sys := newFakeSys() + userSvc := NewUserService(db, sys, testConfig()) + adminSvc := NewAdminService(db) + auditSvc := NewAuditService(db) + cfg := testConfig() + + mailer := &recordingMailer{} + otps := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin) + captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL) + sessions := auth.NewDBSessionStore(db) + resets := auth.NewDBResetTokenStore(db) + limiter := auth.NewRateLimiter(cfg.Auth.MaxLoginFailures, cfg.Auth.LockDuration) + + svc := NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, testLogger()) + return svc, sys, mailer +} + +func TestAuthServiceAdminLogin(t *testing.T) { + svc, _, _ := newTestAuthService(t) + ctx := context.Background() + mustAdmin(t, NewAdminService(svc.db)) + + sid, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test") + if err != nil { + t.Fatalf("admin login: %v", err) + } + if sid == "" { + t.Fatal("session id should not be empty") + } + // 错误密码 + if _, err := svc.AdminLogin(ctx, "root", "wrongpass", "127.0.0.1", "test"); !errors.Is(err, ErrBadCredentials) { + t.Fatalf("bad password err = %v, want ErrBadCredentials", err) + } +} + +func TestAuthServiceAdminLoginRateLimited(t *testing.T) { + svc, _, _ := newTestAuthService(t) + ctx := context.Background() + mustAdmin(t, NewAdminService(svc.db)) + + for i := 0; i < 5; i++ { + _, _ = svc.AdminLogin(ctx, "root", "wrongpass", "127.0.0.1", "test") + } + _, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test") + if !errors.Is(err, ErrRateLimited) { + t.Fatalf("login after lock err = %v, want ErrRateLimited", err) + } +} + +func TestAuthServiceUserOTPFlow(t *testing.T) { + svc, sys, mailer := newTestAuthService(t) + ctx := context.Background() + + // 准备一个外部用户 + userSvc := NewUserService(svc.db, sys, testConfig()) + u, err := userSvc.Create(ctx, "zhangsan", "zs@example.com", "", "", 0, 0) + if err != nil { + t.Fatalf("create user: %v", err) + } + + // 获取图形验证码 + cap, err := svc.captchas.New() + if err != nil { + t.Fatalf("captcha: %v", err) + } + + // 发送 OTP + if err := svc.UserOTPSend(ctx, u.Username, cap.ID, cap.Text, "127.0.0.1"); err != nil { + t.Fatalf("otp send: %v", err) + } + if !strings.Contains(mailer.lastBody, "验证码") { + t.Fatalf("otp mail body unexpected: %s", mailer.lastBody) + } + + // CLI 通道复用同一验证码 + code, err := svc.otps.Current(ctx, u.Username) + if err != nil { + t.Fatalf("otp current: %v", err) + } + if len(code) != 6 { + t.Fatalf("otp len = %d, want 6", len(code)) + } + + // OTP 登录 + sid, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test") + if err != nil { + t.Fatalf("otp login: %v", err) + } + if sid == "" { + t.Fatal("session id should not be empty") + } + + // 一次性:再次使用失败 + if _, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test"); err != auth.ErrInvalidCode { + t.Fatalf("reuse otp err = %v, want ErrInvalidCode", err) + } + + // 验证码错误时 send 被拒绝 + cap2, _ := svc.captchas.New() + if err := svc.UserOTPSend(ctx, u.Username, cap2.ID, "0000", "127.0.0.1"); !errors.Is(err, ErrCaptchaFailed) { + t.Fatalf("bad captcha err = %v, want ErrCaptchaFailed", err) + } + + // 冷却期内再次 send 被拒 + cap3, _ := svc.captchas.New() + if err := svc.UserOTPSend(ctx, u.Username, cap3.ID, cap3.Text, "127.0.0.1"); !errors.Is(err, auth.ErrCooldown) { + t.Fatalf("cooldown err = %v, want ErrCooldown", err) + } +} + +func TestAuthServiceAdminForgotReset(t *testing.T) { + svc, _, mailer := newTestAuthService(t) + ctx := context.Background() + adm := mustAdmin(t, NewAdminService(svc.db)) + + // forgot → 邮件应包含重置链接 + if err := svc.AdminForgot(ctx, adm.Username, "127.0.0.1"); err != nil { + t.Fatalf("forgot: %v", err) + } + if mailer.lastTo != adm.Email { + t.Fatalf("mail to = %q, want %q", mailer.lastTo, adm.Email) + } + // 从邮件 body 提取 token + idx := strings.Index(mailer.lastBody, "token=") + if idx < 0 { + t.Fatalf("reset link missing token: %s", mailer.lastBody) + } + token := mailer.lastBody[idx+len("token="):] + token = strings.TrimSpace(strings.SplitN(token, "\n", 2)[0]) + + // reset + if err := svc.AdminReset(ctx, token, "NewPassw0rd", "127.0.0.1"); err != nil { + t.Fatalf("reset: %v", err) + } + // 新密码可登录 + adminSvc := NewAdminService(svc.db) + if _, err := adminSvc.Login(ctx, adm.Username, "NewPassw0rd"); err != nil { + t.Fatalf("login with new password: %v", err) + } + // 旧密码失效 + if _, err := adminSvc.Login(ctx, adm.Username, "Passw0rd"); !errors.Is(err, ErrBadCredentials) { + t.Fatalf("login with old password err = %v, want ErrBadCredentials", err) + } + // 令牌一次性 + if err := svc.AdminReset(ctx, token, "AgainPassw0rd", "127.0.0.1"); !errors.Is(err, auth.ErrResetTokenInvalid) { + t.Fatalf("reuse token err = %v, want ErrResetTokenInvalid", err) + } + // 不存在的用户 forgot 也应成功(防枚举) + if err := svc.AdminForgot(ctx, "ghost", "127.0.0.1"); err != nil { + t.Fatalf("forgot ghost: %v", err) + } +} + +func TestAuthServiceMeLogout(t *testing.T) { + svc, sys, _ := newTestAuthService(t) + ctx := context.Background() + mustAdmin(t, NewAdminService(svc.db)) + + sid, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test") + if err != nil { + t.Fatalf("login: %v", err) + } + info, err := svc.Me(ctx, sid) + if err != nil { + t.Fatalf("me: %v", err) + } + if info.UserType != "admin" || info.Username != "root" { + t.Fatalf("me info = %+v", info) + } + if err := svc.Logout(ctx, sid); err != nil { + t.Fatalf("logout: %v", err) + } + if _, err := svc.Me(ctx, sid); !errors.Is(err, auth.ErrSessionNotFound) { + t.Fatalf("me after logout err = %v, want ErrSessionNotFound", err) + } + + // 外部用户 me + userSvc := NewUserService(svc.db, sys, testConfig()) + u, _ := userSvc.Create(ctx, "lisi", "ls@example.com", "", "", 0, 0) + cap, _ := svc.captchas.New() + _ = svc.UserOTPSend(ctx, u.Username, cap.ID, cap.Text, "127.0.0.1") + code, _ := svc.otps.Current(ctx, u.Username) + usid, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test") + if err != nil { + t.Fatalf("user login: %v", err) + } + info, err = svc.Me(ctx, usid) + if err != nil { + t.Fatalf("user me: %v", err) + } + if info.UserType != "user" || info.Username != "ext_lisi" { + t.Fatalf("user me info = %+v", info) + } +} diff --git a/internal/service/user.go b/internal/service/user.go index 50420fe..daeac1c 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -4,9 +4,11 @@ import ( "context" "errors" "strings" + "time" "gorm.io/gorm" + "ws_usernode/internal/config" "ws_usernode/internal/model" "ws_usernode/internal/pkg" "ws_usernode/internal/system" @@ -14,29 +16,30 @@ import ( // 用户服务错误。 var ( - ErrUserNotFound = errors.New("service: 用户不存在") - ErrUserExists = errors.New("service: 用户名已存在") + ErrUserNotFound = errors.New("service: 用户不存在") + ErrUserExists = errors.New("service: 用户名已存在") + ErrUserExpired = errors.New("service: 用户已过期,请先延期") + ErrUserDisabled = errors.New("service: 用户已禁用,无法操作") + ErrSystemAccountMissing = errors.New("service: 系统账号不存在,无法操作") ) -// UserService 外部用户生命周期服务。 -// M0 提供查询与创建骨架;建号(useradd)、禁用/延期等系统操作 M1 接入 system.Manager。 +// UserService 外部用户生命周期服务:DB 记录 + system.Manager 系统账号操作。 type UserService struct { db *gorm.DB sys system.Manager + cfg *config.Config } // NewUserService 创建用户服务。 -func NewUserService(db *gorm.DB, sys system.Manager) *UserService { - return &UserService{db: db, sys: sys} +func NewUserService(db *gorm.DB, sys system.Manager, cfg *config.Config) *UserService { + return &UserService{db: db, sys: sys, cfg: cfg} } // GetByUsername 按用户名查询外部用户(含或不含 ext_ 前缀均可)。 func (s *UserService) GetByUsername(ctx context.Context, username string) (*model.User, error) { - if !strings.HasPrefix(username, "ext_") { - username = "ext_" + username - } + full := normalizeName(username, s.cfg.System.UserPrefix) var u model.User - if err := s.db.WithContext(ctx).First(&u, "username = ?", username).Error; err != nil { + if err := s.db.WithContext(ctx).First(&u, "username = ?", full).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, ErrUserNotFound } @@ -45,16 +48,68 @@ func (s *UserService) GetByUsername(ctx context.Context, username string) (*mode return &u, nil } -// Create 创建外部用户记录并调用系统层建号。M0 阶段系统层为 dry-run。 -// username 为不含前缀的申请名,内部加 ext_ 前缀。 -func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttlSeconds int64) (*model.User, error) { +// GetByID 按 ID 查询外部用户。 +func (s *UserService) GetByID(ctx context.Context, id uint) (*model.User, error) { + var u model.User + if err := s.db.WithContext(ctx).First(&u, "id = ?", id).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrUserNotFound + } + return nil, err + } + return &u, nil +} + +// UserFilter 用户列表筛选条件。 +type UserFilter struct { + Status string // active / disabled / expired,空为全部 + Supervisor string // 挂靠老师模糊匹配 + Page int + PageSize int +} + +// List 分页查询用户(admin)。 +func (s *UserService) List(ctx context.Context, f UserFilter) ([]model.User, int64, error) { + q := s.db.WithContext(ctx).Model(&model.User{}) + if f.Status != "" { + q = q.Where("status = ?", f.Status) + } + if f.Supervisor != "" { + q = q.Where("supervisor LIKE ?", "%"+f.Supervisor+"%") + } + 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 > 100 { + size = 100 + } + var users []model.User + if err := q.Order("id DESC").Offset((page - 1) * size).Limit(size).Find(&users).Error; err != nil { + return nil, 0, err + } + return users, total, nil +} + +// Create 创建外部用户:DB 记录 + 系统账号(useradd + passwd -l)。 +// username 不含前缀;ttl 为有效期时长,0 表示用配置默认(90 天)。 +// 系统建号失败时回滚 DB 记录,保证两侧一致。 +func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttl time.Duration, createdBy uint) (*model.User, error) { if err := pkg.ValidateUserName(username); err != nil { return nil, err } if err := pkg.ValidateEmail(email); err != nil { return nil, err } - full := "ext_" + username + username = strings.TrimSpace(username) + full := s.cfg.System.UserPrefix + username var count int64 if err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", full).Count(&count).Error; err != nil { return nil, err @@ -62,21 +117,178 @@ func (s *UserService) Create(ctx context.Context, username, email, supervisor, p if count > 0 { return nil, ErrUserExists } + if ttl <= 0 { + ttl = s.cfg.Policy.DefaultTTL + } + expireAt := time.Now().Add(ttl) u := &model.User{ Username: full, Email: email, Supervisor: supervisor, Purpose: purpose, Status: model.UserStatusActive, - Shell: "/bin/sh", + ExpireAt: &expireAt, + Shell: s.cfg.System.Shell, + CreatedBy: createdBy, } if err := s.db.WithContext(ctx).Create(u).Error; err != nil { return nil, err } - // 系统账号创建(dry-run / 真实),失败时回滚 DB 记录 - if err := s.sys.CreateUser(ctx, system.Account{Username: full}); err != nil { + // 系统账号创建(dry-run / 直接 / sudo),失败时回滚 DB 记录 + if err := s.sys.CreateUser(ctx, system.Account{Username: full, Shell: s.cfg.System.Shell}); err != nil { _ = s.db.WithContext(ctx).Delete(u).Error return nil, err } return u, nil } + +// Update 更新外部用户信息(仅更新非 nil 字段;邮箱由管理员修改,用户不可自助改)。 +func (s *UserService) Update(ctx context.Context, id uint, email, supervisor, purpose *string) (*model.User, error) { + if _, err := s.GetByID(ctx, id); err != nil { + return nil, err + } + updates := make(map[string]any) + if email != nil { + if err := pkg.ValidateEmail(*email); err != nil { + return nil, err + } + updates["email"] = *email + } + if supervisor != nil { + updates["supervisor"] = *supervisor + } + if purpose != nil { + updates["purpose"] = *purpose + } + if len(updates) > 0 { + if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error; err != nil { + return nil, err + } + } + return s.GetByID(ctx, id) +} + +// systemAccountOK 检查系统账号存在;dry-run 模式跳过检查(演练流程)。 +func (s *UserService) systemAccountOK(ctx context.Context, username string) bool { + if s.cfg.System.DryRun { + return true + } + ok, err := s.sys.Exists(ctx, username) + return err == nil && ok +} + +// Disable 禁用用户:DB 置 disabled + 清空 authorized_keys(SSH 立即失效)。 +func (s *UserService) Disable(ctx context.Context, id uint) error { + u, err := s.GetByID(ctx, id) + if err != nil { + return err + } + if u.Status == model.UserStatusDisabled { + return nil // 幂等 + } + if !s.systemAccountOK(ctx, u.Username) { + return ErrSystemAccountMissing + } + if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, nil); err != nil { + return err + } + return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Update("status", model.UserStatusDisabled).Error +} + +// Enable 启用用户:DB 置 active + 按 DB 状态重写 authorized_keys。 +// 已过期的用户需先延期(Extend)。 +func (s *UserService) Enable(ctx context.Context, id uint) error { + u, err := s.GetByID(ctx, id) + if err != nil { + return err + } + if u.Status == model.UserStatusActive { + return nil // 幂等 + } + if u.ExpireAt != nil && time.Now().After(*u.ExpireAt) { + return ErrUserExpired + } + if !s.systemAccountOK(ctx, u.Username) { + return ErrSystemAccountMissing + } + // 恢复有效密钥(M1 阶段用户尚无密钥,M2 接入后按 DB 同步) + keys, err := s.activeKeys(ctx, u.ID) + if err != nil { + return err + } + if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, keys); err != nil { + return err + } + return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Update("status", model.UserStatusActive).Error +} + +// Extend 延长有效期:重设 expire_at(days<=0 用配置默认 TTL)。 +// 已过期用户在回收期内可经此恢复(PLAN §2.2),恢复后同步密钥。 +func (s *UserService) Extend(ctx context.Context, id uint, days int) error { + u, err := s.GetByID(ctx, id) + if err != nil { + return err + } + ttl := time.Duration(days) * 24 * time.Hour + if days <= 0 { + ttl = s.cfg.Policy.DefaultTTL + } + 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 + } + keys, err := s.activeKeys(ctx, u.ID) + if err != nil { + return err + } + if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, keys); err != nil { + return err + } + updates["status"] = model.UserStatusActive + } + return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error +} + +// Delete 删除并回收用户:删除系统账号(userdel -r)+ 家目录 + 密钥记录, +// 保留审计。系统账号已不存在时仍完成 DB 清理。 +func (s *UserService) Delete(ctx context.Context, id uint) error { + u, err := s.GetByID(ctx, id) + if err != nil { + return err + } + if s.systemAccountOK(ctx, u.Username) { + if err := s.sys.RemoveUser(ctx, u.Username); err != nil { + return err + } + } + return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("user_id = ?", u.ID).Delete(&model.SSHKey{}).Error; err != nil { + return err + } + return tx.Delete(&model.User{}, "id = ?", u.ID).Error + }) +} + +// activeKeys 返回用户当前有效(active)密钥,供授权同步(M2 完善密钥管理)。 +func (s *UserService) activeKeys(ctx context.Context, userID uint) ([]system.Key, error) { + var rows []model.SSHKey + if err := s.db.WithContext(ctx).Where("user_id = ? AND status = ?", userID, model.StatusActive).Find(&rows).Error; err != nil { + return nil, err + } + keys := make([]system.Key, 0, len(rows)) + for _, k := range rows { + keys = append(keys, system.Key{Type: k.KeyType, PublicKey: k.PublicKey}) + } + return keys, nil +} + +// normalizeName 补全系统账号前缀(如 ext_)。 +func normalizeName(username, prefix string) string { + name := strings.TrimSpace(username) + if !strings.HasPrefix(name, prefix) { + return prefix + name + } + return name +} diff --git a/internal/service/user_test.go b/internal/service/user_test.go new file mode 100644 index 0000000..119f3db --- /dev/null +++ b/internal/service/user_test.go @@ -0,0 +1,242 @@ +package service + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "gorm.io/gorm" + + "ws_usernode/internal/config" + "ws_usernode/internal/model" + "ws_usernode/internal/system" +) + +// fakeSys 内存版 system.Manager:记录已创建的账号,便于断言与真实系统隔离。 +type fakeSys struct { + mu sync.Mutex + accounts map[string]bool + keys map[string][]system.Key // username -> 最后一次同步的密钥 + lastCmd string +} + +func newFakeSys() *fakeSys { + return &fakeSys{accounts: map[string]bool{}, keys: map[string][]system.Key{}} +} + +func (f *fakeSys) CreateUser(_ context.Context, acc system.Account) error { + f.mu.Lock() + defer f.mu.Unlock() + f.accounts[acc.Username] = true + f.lastCmd = "useradd " + acc.Username + return nil +} + +func (f *fakeSys) RemoveUser(_ context.Context, username string) error { + f.mu.Lock() + defer f.mu.Unlock() + delete(f.accounts, username) + f.lastCmd = "userdel " + username + return nil +} + +func (f *fakeSys) SetLock(_ context.Context, username string, locked bool) error { + f.mu.Lock() + defer f.mu.Unlock() + f.lastCmd = "passwd " + username + return nil +} + +func (f *fakeSys) Exists(_ context.Context, username string) (bool, error) { + f.mu.Lock() + defer f.mu.Unlock() + return f.accounts[username], nil +} + +func (f *fakeSys) SyncAuthorizedKeys(_ context.Context, username string, keys []system.Key) error { + f.mu.Lock() + defer f.mu.Unlock() + f.keys[username] = keys + f.lastCmd = "sync-keys " + username + return nil +} + +func (f *fakeSys) has(username string) bool { + f.mu.Lock() + defer f.mu.Unlock() + return f.accounts[username] +} + +func testDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := model.Open("sqlite", ":memory:", false) + if err != nil { + t.Fatalf("open test db: %v", err) + } + if err := model.Migrate(db); err != nil { + t.Fatalf("migrate: %v", err) + } + return db +} + +func testConfig() *config.Config { + cfg := config.Default() + cfg.System.DryRun = false // 测试直接走 fakeSys,不依赖 dry-run + return cfg +} + +func mustAdmin(t *testing.T, svc *AdminService) *model.AdminUser { + t.Helper() + adm, err := svc.Create(context.Background(), "root", "Passw0rd", "root@example.com") + if err != nil { + t.Fatalf("create admin: %v", err) + } + return adm +} + +func TestAdminServiceLogin(t *testing.T) { + db := testDB(t) + svc := NewAdminService(db) + adm := mustAdmin(t, svc) + + got, err := svc.Login(context.Background(), adm.Username, "Passw0rd") + if err != nil { + t.Fatalf("login: %v", err) + } + if got.ID != adm.ID { + t.Fatalf("login returned wrong admin") + } + if _, err := svc.Login(context.Background(), adm.Username, "wrong"); !errors.Is(err, ErrBadCredentials) { + t.Fatalf("bad password err = %v, want ErrBadCredentials", err) + } + if _, err := svc.Login(context.Background(), "ghost", "Passw0rd"); !errors.Is(err, ErrBadCredentials) { + t.Fatalf("missing user err = %v, want ErrBadCredentials", err) + } +} + +func TestUserServiceLifecycle(t *testing.T) { + db := testDB(t) + sys := newFakeSys() + svc := NewUserService(db, sys, testConfig()) + ctx := context.Background() + + u, err := svc.Create(ctx, "zhangsan", "zs@example.com", "prof.li", "科研", 90*24*time.Hour, 1) + if err != nil { + t.Fatalf("create: %v", err) + } + if u.Username != "ext_zhangsan" { + t.Fatalf("username = %q, want ext_zhangsan", u.Username) + } + if !sys.has("ext_zhangsan") { + t.Fatal("system account should exist after create") + } + if u.ExpireAt == nil { + t.Fatal("expire_at should be set") + } + + // 重复创建冲突 + if _, err := svc.Create(ctx, "zhangsan", "x@example.com", "", "", 0, 1); !errors.Is(err, ErrUserExists) { + t.Fatalf("duplicate create err = %v, want ErrUserExists", err) + } + + // 列表 + users, total, err := svc.List(ctx, UserFilter{Page: 1, PageSize: 20}) + if err != nil { + t.Fatalf("list: %v", err) + } + if total != 1 || len(users) != 1 { + t.Fatalf("list total=%d len=%d, want 1/1", total, len(users)) + } + + // 更新 + email := "new@example.com" + supp := "prof.wang" + u2, err := svc.Update(ctx, u.ID, &email, &supp, nil) + if err != nil { + t.Fatalf("update: %v", err) + } + if u2.Email != "new@example.com" || u2.Supervisor != "prof.wang" { + t.Fatalf("update not applied: %+v", u2) + } + + // 禁用 → 系统密钥清空 + if err := svc.Disable(ctx, u.ID); err != nil { + t.Fatalf("disable: %v", err) + } + if u2, _ := svc.GetByID(ctx, u.ID); u2.Status != model.UserStatusDisabled { + t.Fatalf("status = %q, want disabled", u2.Status) + } + + // 启用 → 恢复 active + if err := svc.Enable(ctx, u.ID); err != nil { + t.Fatalf("enable: %v", err) + } + if u2, _ := svc.GetByID(ctx, u.ID); u2.Status != model.UserStatusActive { + t.Fatalf("status = %q, want active", u2.Status) + } + + // 延期 + if err := svc.Extend(ctx, u.ID, 30); err != nil { + t.Fatalf("extend: %v", err) + } + + // 删除 → 系统账号移除 + if err := svc.Delete(ctx, u.ID); err != nil { + t.Fatalf("delete: %v", err) + } + if sys.has("ext_zhangsan") { + t.Fatal("system account should be removed after delete") + } + if _, err := svc.GetByID(ctx, u.ID); !errors.Is(err, ErrUserNotFound) { + t.Fatalf("get after delete err = %v, want ErrUserNotFound", err) + } +} + +func TestUserServiceEnableExpired(t *testing.T) { + db := testDB(t) + sys := newFakeSys() + svc := NewUserService(db, sys, testConfig()) + ctx := context.Background() + + u, err := svc.Create(ctx, "lisi", "ls@example.com", "", "", 0, 1) + if err != nil { + t.Fatalf("create: %v", err) + } + // 强制置为过期 + past := time.Now().Add(-time.Hour) + if err := db.Model(&model.User{}).Where("id = ?", u.ID).Update("expire_at", &past).Error; err != nil { + t.Fatalf("force expire: %v", err) + } + db.Model(&model.User{}).Where("id = ?", u.ID).Update("status", model.UserStatusExpired) + + if err := svc.Enable(ctx, u.ID); !errors.Is(err, ErrUserExpired) { + t.Fatalf("enable expired err = %v, want ErrUserExpired", err) + } + // 延期可恢复 + 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 { + t.Fatalf("status after extend = %q, want active", u2.Status) + } +} + +func TestUserServiceSystemAccountMissing(t *testing.T) { + db := testDB(t) + sys := newFakeSys() + cfg := testConfig() + cfg.System.DryRun = false + svc := NewUserService(db, sys, cfg) + ctx := context.Background() + + // 手动插一条 DB 记录,但系统账号不存在 + u := &model.User{Username: "ext_orphan", Email: "o@example.com", Status: model.UserStatusActive, Shell: "/bin/sh"} + if err := db.Create(u).Error; err != nil { + t.Fatalf("insert: %v", err) + } + if err := svc.Disable(ctx, u.ID); !errors.Is(err, ErrSystemAccountMissing) { + t.Fatalf("disable missing account err = %v, want ErrSystemAccountMissing", err) + } +} diff --git a/internal/system/system.go b/internal/system/system.go index 21fc8cb..4f6c2f4 100644 --- a/internal/system/system.go +++ b/internal/system/system.go @@ -4,15 +4,19 @@ // 执行 useradd/usermod/userdel/passwd 等固定命令并做参数强校验。未来多节点 // agent 模式只需新增远程实现替换本地实现(PLAN §5.1)。 // -// 权限模型:节点以专有用户(如 usernode)运行,经 sudo -n 提权执行白名单 -// 命令;开发环境 config system.sudo=false 进入 dry-run(只打印计划不执行), -// 避免在开发机上直接操作系统账号。系统命令的真实系统效果在 M1 用测试用户/ -// 容器验证,不在生产直接跑 useradd。 +// 执行模式(config system.*): +// - dry_run=true:只打印计划命令不执行(开发演练); +// - dry_run=false + sudo=false:直接执行(容器/测试用户验证真实建号); +// - dry_run=false + sudo=true:经 sudo -n 提权执行白名单命令(生产, +// 需配置 deploy/sudoers)。真实系统账号操作只在测试用户/容器中验证, +// 不在生产直接跑 useradd。 package system import ( + "bufio" "context" "errors" + "fmt" "os" "os/exec" "path/filepath" @@ -45,8 +49,10 @@ type Manager interface { RemoveUser(ctx context.Context, username string) error // SetLock 锁定/解锁系统账号口令(passwd -l / -u)。 SetLock(ctx context.Context, username string, locked bool) error + // Exists 检查系统账号是否存在(读 /etc/passwd,无需提权)。 + Exists(ctx context.Context, username string) (bool, error) // SyncAuthorizedKeys 以 DB 状态全量重写 authorized_keys(原子写 + 并发锁), - // 吊销密钥即从文件移除、立即失效。M0 提供 dry-run 实现,M2 完成生产路径。 + // 吊销密钥即从文件移除、立即失效。 SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error } @@ -69,20 +75,32 @@ func New(cfg config.SystemConfig) Manager { return &localManager{cfg: cfg} } -// run 执行白名单命令:sudo -n ,参数在调用处强校验。 -// dry-run 模式下只返回将执行的命令文本,不真正执行。 +// run 执行白名单命令。参数在调用处强校验。 +// 模式:dry-run 只返回命令文本;direct(sudo=false)直接执行; +// sudo 经 sudo -n 提权执行(对应 deploy/sudoers 白名单)。 func (m *localManager) run(ctx context.Context, cmd string, args ...string) (string, error) { if !allowedCommands[cmd] { return "", errors.New("system: command not allowed: " + cmd) } - argv := append([]string{"-n", cmd}, args...) - cmdline := strings.Join(append([]string{"sudo", "-n", cmd}, args...), " ") - if !m.cfg.Sudo { - return cmdline, nil // dry-run + cmdline := strings.Join(append([]string{cmd}, args...), " ") + if m.cfg.DryRun { + return cmdline, nil + } + var ( + out []byte + err error + ) + if m.cfg.Sudo { + cmdline = "sudo -n " + cmdline + out, err = exec.CommandContext(ctx, "sudo", append([]string{"-n", cmd}, args...)...).CombinedOutput() + } else { + out, err = exec.CommandContext(ctx, cmd, args...).CombinedOutput() } - out, err := exec.CommandContext(ctx, "sudo", argv...).CombinedOutput() if err != nil { - return string(out), err + if len(out) > 0 { + return cmdline, fmt.Errorf("%s: %s: %w", cmdline, strings.TrimSpace(string(out)), err) + } + return cmdline, fmt.Errorf("%s: %w", cmdline, err) } return cmdline, nil } @@ -109,11 +127,11 @@ func (m *localManager) CreateUser(ctx context.Context, acc Account) error { home = filepath.Join(m.cfg.HomeBase, username) } if _, err := m.run(ctx, "useradd", "-m", "-d", home, "-s", shell, "-g", m.cfg.Group, username); err != nil { - return errors.New("system: useradd: " + err.Error()) + return err } // 锁定口令,仅密钥登录 if _, err := m.run(ctx, "passwd", "-l", username); err != nil { - return errors.New("system: passwd -l: " + err.Error()) + return err } return nil } @@ -124,7 +142,7 @@ func (m *localManager) RemoveUser(ctx context.Context, username string) error { return err } if _, err := m.run(ctx, "userdel", "-r", username); err != nil { - return errors.New("system: userdel: " + err.Error()) + return err } return nil } @@ -139,20 +157,51 @@ func (m *localManager) SetLock(ctx context.Context, username string, locked bool flag = "-l" } if _, err := m.run(ctx, "passwd", flag, username); err != nil { - return errors.New("system: passwd: " + err.Error()) + return err } return nil } +// passwdPath 为系统账号数据库路径(测试可覆盖为临时文件)。 +var passwdPath = "/etc/passwd" + +// Exists 检查系统账号是否存在。直接读 /etc/passwd(世界可读,无需提权)。 +func (m *localManager) Exists(ctx context.Context, username string) (bool, error) { + username = m.sysName(username) + if err := pkg.ValidateSystemAccount(username); err != nil { + return false, err + } + _ = ctx + f, err := os.Open(passwdPath) + if err != nil { + return false, err + } + defer f.Close() + sc := bufio.NewScanner(f) + for sc.Scan() { + line := sc.Text() + if line == "" { + continue + } + fields := strings.SplitN(line, ":", 2) + if len(fields) > 0 && fields[0] == username { + return true, nil + } + } + if err := sc.Err(); err != nil { + return false, err + } + return false, nil +} + // SyncAuthorizedKeys 全量重写 authorized_keys: // // 1. 以 /.lock 文件锁串行化并发写(flock); // 2. 写临时文件(0600),再 rename 原子替换; // 3. 无有效密钥时写入空文件(SSH 行为一致,避免文件缺失导致的歧义)。 // -// 生产模式(sudo=true)下节点进程需具备对家目录 .ssh 的写权限——部署时经 -// sudoers 白名单授予固定命令/受控脚本实现(M2 细化);当前实现直接做文件 -// 操作并假设权限已配置,dry-run 模式打印计划命令。 +// dry-run 模式只打印计划;真实写路径(direct/sudo)需要进程具备对家目录 +// .ssh 的写权限——生产部署经 sudoers 白名单授予受控脚本实现(M2 细化)。 func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error { username = m.sysName(username) if err := pkg.ValidateSystemAccount(username); err != nil { @@ -169,13 +218,13 @@ func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, } } - if !m.cfg.Sudo { + if m.cfg.DryRun { var plan strings.Builder plan.WriteString("mkdir -p " + sshDir + " (0700)\n") plan.WriteString("flock " + lockPath + "\n") plan.WriteString("write " + filepath.Join(sshDir, "authorized_keys") + " (0600)\n") plan.WriteString(content.String()) - // 骨架阶段:仅日志输出计划,不落盘 + // 演练模式:仅日志输出计划,不落盘 return nil } diff --git a/internal/system/system_test.go b/internal/system/system_test.go new file mode 100644 index 0000000..18a1982 --- /dev/null +++ b/internal/system/system_test.go @@ -0,0 +1,87 @@ +package system + +import ( + "context" + "os" + "path/filepath" + "testing" + + "ws_usernode/internal/config" +) + +func testCfg() config.SystemConfig { + cfg := config.Default() + return cfg.System +} + +func TestDryRunDoesNotExecute(t *testing.T) { + cfg := testCfg() + cfg.DryRun = true + m := New(cfg) + ctx := context.Background() + + // dry-run 下创建/删除不应报错(只打印计划) + if err := m.CreateUser(ctx, Account{Username: "ext_zhangsan"}); err != nil { + t.Fatalf("dry-run create: %v", err) + } + if err := m.RemoveUser(ctx, "ext_zhangsan"); err != nil { + t.Fatalf("dry-run remove: %v", err) + } + if err := m.SetLock(ctx, "ext_zhangsan", true); err != nil { + t.Fatalf("dry-run lock: %v", err) + } + if err := m.SyncAuthorizedKeys(ctx, "ext_zhangsan", nil); err != nil { + t.Fatalf("dry-run keys: %v", err) + } + // 参数校验在 dry-run 下仍然生效 + if err := m.CreateUser(ctx, Account{Username: "ext_..bad"}); err == nil { + t.Fatal("expected validation error for illegal account name") + } +} + +func TestExistsParsesPasswd(t *testing.T) { + // 用临时 passwd 文件验证解析逻辑 + dir := t.TempDir() + p := filepath.Join(dir, "passwd") + content := "root:x:0:0:root:/root:/bin/sh\n" + + "ext_zhangsan:x:1001:1001::/home/ext_zhangsan:/bin/sh\n" + if err := os.WriteFile(p, []byte(content), 0o600); err != nil { + t.Fatal(err) + } + old := passwdPath + passwdPath = p + t.Cleanup(func() { passwdPath = old }) + + m := New(testCfg()) + ctx := context.Background() + // 已带前缀与未带前缀都会命中同一账号 + for _, name := range []string{"ext_zhangsan", "zhangsan"} { + ok, err := m.Exists(ctx, name) + if err != nil { + t.Fatalf("exists(%q): %v", name, err) + } + if !ok { + t.Fatalf("user %q should exist", name) + } + } + // 不存在的账号返回 false + if ok, err := m.Exists(ctx, "ghost_xyz"); err != nil || ok { + t.Fatalf("exists(ghost) = %v/%v, want false/nil", ok, err) + } + // 非法账号名直接报错 + if _, err := m.Exists(ctx, "bad..name"); err == nil { + t.Fatal("expected validation error") + } +} + +func TestPrefixNormalization(t *testing.T) { + cfg := testCfg() + lm := &localManager{cfg: cfg} + // sysName 逻辑:已带前缀不重复加 + if got := lm.sysName("ext_x"); got != "ext_x" { + t.Fatalf("sysName(ext_x) = %q", got) + } + if got := lm.sysName("x"); got != "ext_x" { + t.Fatalf("sysName(x) = %q", got) + } +}