feat(M1): 认证与用户管理 — 双通道登录、cookie 会话、用户 CRUD 与真实系统账号对接
认证: - 图形验证码 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 集成 + 容器内真实系统账号端到端验证
This commit is contained in:
@@ -13,7 +13,7 @@ NET_HOST := --network=host
|
|||||||
GO ?= go
|
GO ?= go
|
||||||
PODMAN ?= podman
|
PODMAN ?= podman
|
||||||
BIN := bin/usernode
|
BIN := bin/usernode
|
||||||
VERSION ?= 0.1.0-m0
|
VERSION ?= 0.2.0-m1
|
||||||
LDFLAGS := -s -w -X main.version=$(VERSION)
|
LDFLAGS := -s -w -X main.version=$(VERSION)
|
||||||
GOFLAGS := -trimpath
|
GOFLAGS := -trimpath
|
||||||
|
|
||||||
|
|||||||
+14
-3
@@ -10,8 +10,10 @@ import (
|
|||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"ws_usernode/internal/api"
|
"ws_usernode/internal/api"
|
||||||
|
"ws_usernode/internal/auth"
|
||||||
"ws_usernode/internal/config"
|
"ws_usernode/internal/config"
|
||||||
"ws_usernode/internal/cron"
|
"ws_usernode/internal/cron"
|
||||||
|
"ws_usernode/internal/mail"
|
||||||
"ws_usernode/internal/model"
|
"ws_usernode/internal/model"
|
||||||
"ws_usernode/internal/router"
|
"ws_usernode/internal/router"
|
||||||
"ws_usernode/internal/server"
|
"ws_usernode/internal/server"
|
||||||
@@ -82,9 +84,18 @@ func cmdServe(args []string) error {
|
|||||||
|
|
||||||
sys := system.New(cfg.System)
|
sys := system.New(cfg.System)
|
||||||
adminSvc := service.NewAdminService(db)
|
adminSvc := service.NewAdminService(db)
|
||||||
userSvc := service.NewUserService(db, sys)
|
userSvc := service.NewUserService(db, sys, cfg)
|
||||||
auditSvc := service.NewAuditService(db)
|
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)
|
// 启动前自动迁移(骨架阶段保证表结构就绪;M5 部署建议显式 migrate)
|
||||||
if err := model.Migrate(db); err != nil {
|
if err := model.Migrate(db); err != nil {
|
||||||
return fmt.Errorf("数据库迁移: %w", err)
|
return fmt.Errorf("数据库迁移: %w", err)
|
||||||
@@ -97,8 +108,8 @@ func cmdServe(args []string) error {
|
|||||||
sched.Start()
|
sched.Start()
|
||||||
defer sched.Stop()
|
defer sched.Stop()
|
||||||
|
|
||||||
h := api.New(adminSvc, userSvc, auditSvc)
|
h := api.New(cfg, authSvc, userSvc, auditSvc)
|
||||||
r := router.New(cfg, h, log)
|
r := router.New(cfg, h, sessions, log)
|
||||||
|
|
||||||
srv := server.New(cfg.Server.Listen, r, log)
|
srv := server.New(cfg.Server.Listen, r, log)
|
||||||
if err := srv.Run(); err != nil {
|
if err := srv.Run(); err != nil {
|
||||||
|
|||||||
+18
-17
@@ -25,8 +25,8 @@ func cmdUser(args []string) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// userOTP 获取外部用户 OTP 验证码,与邮件通道共用同一存储与限速
|
// userOTP 获取外部用户 OTP 验证码。与邮件通道共用同一 DB 存储与限速:
|
||||||
// (同一验证码、同一 10 分钟有效期、同一 60s 冷却与失败限速)。
|
// 已有有效验证码时直接复用(同一验证码),无则生成(受同一 60s 冷却约束)。
|
||||||
func userOTP(args []string) error {
|
func userOTP(args []string) error {
|
||||||
fs := flag.NewFlagSet("user otp", flag.ContinueOnError)
|
fs := flag.NewFlagSet("user otp", flag.ContinueOnError)
|
||||||
cfgPath, debug := commonFlags(fs)
|
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)
|
return fmt.Errorf("用户不存在或不可用: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// M0 骨架:CLI 独立生成(内存 store 与运行中服务不共享)。
|
store := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin)
|
||||||
// 生产对齐(同一验证码/冷却/限速跨通道生效)需 OTP 落 DB,M1 实现
|
ctx := context.Background()
|
||||||
// auth.OTPStore 的 DB 实现后,CLI 与邮件通道读写同一存储。
|
code, err := store.Current(ctx, name)
|
||||||
store := auth.NewMemoryOTPStore()
|
|
||||||
code, err := store.Send(name, cfg.Policy.OTPTTL, cfg.Policy.OTPCooldown)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if err == auth.ErrCooldown {
|
// 无有效验证码:生成(先到先得,覆盖旧码;冷却期内返回 ErrCooldown)
|
||||||
return fmt.Errorf("发送冷却中,请稍后重试(冷却 %s)", cfg.Policy.OTPCooldown)
|
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("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
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
+8
-1
@@ -8,6 +8,7 @@
|
|||||||
[app]
|
[app]
|
||||||
name = "ws_usernode"
|
name = "ws_usernode"
|
||||||
env = "development" # development / production
|
env = "development" # development / production
|
||||||
|
base_url = "http://127.0.0.1:8080" # 对外访问地址(邮件重置链接等)
|
||||||
|
|
||||||
[server]
|
[server]
|
||||||
listen = "127.0.0.1:8080" # 生产建议 0.0.0.0:8080 并置于反向代理后
|
listen = "127.0.0.1:8080" # 生产建议 0.0.0.0:8080 并置于反向代理后
|
||||||
@@ -30,6 +31,11 @@ audit_retention = "720h" # 审计保留 30 天,保留前先归档
|
|||||||
otp_ttl = "10m" # OTP 验证码有效期
|
otp_ttl = "10m" # OTP 验证码有效期
|
||||||
otp_cooldown = "60s" # OTP 发送冷却
|
otp_cooldown = "60s" # OTP 发送冷却
|
||||||
|
|
||||||
|
[auth]
|
||||||
|
max_login_failures = 5 # 管理员登录连续失败阈值,达到后锁定
|
||||||
|
lock_duration = "15m" # 锁定持续时间
|
||||||
|
captcha_ttl = "5m" # 图形验证码有效期
|
||||||
|
|
||||||
[smtp]
|
[smtp]
|
||||||
host = "" # 留空则禁用邮件(OTP 仍可用 CLI 通道获取)
|
host = "" # 留空则禁用邮件(OTP 仍可用 CLI 通道获取)
|
||||||
port = 587
|
port = 587
|
||||||
@@ -38,7 +44,8 @@ password = ""
|
|||||||
from = "usernode@example.com"
|
from = "usernode@example.com"
|
||||||
|
|
||||||
[system]
|
[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_" # 外部用户系统账号统一前缀
|
user_prefix = "ext_" # 外部用户系统账号统一前缀
|
||||||
group = "external" # 外部用户统一组
|
group = "external" # 外部用户统一组
|
||||||
shell = "/bin/sh" # 默认 shell
|
shell = "/bin/sh" # 默认 shell
|
||||||
|
|||||||
@@ -49,7 +49,8 @@ RUN CGO_ENABLED=0 go build -trimpath -ldflags="-s -w" -o /out/usernode ./cmd/use
|
|||||||
|
|
||||||
# ---------- 阶段 3:运行镜像 ----------
|
# ---------- 阶段 3:运行镜像 ----------
|
||||||
FROM alpine:3.20
|
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
|
&& addgroup -S usernode && adduser -S -G usernode usernode
|
||||||
COPY --from=go-build /out/usernode /usr/local/bin/usernode
|
COPY --from=go-build /out/usernode /usr/local/bin/usernode
|
||||||
|
|
||||||
|
|||||||
@@ -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 <home> -s <shell> -g external <name> 创建账号
|
||||||
|
# - usermod 预留(如 usermod -e 过期),M4 回收期使用
|
||||||
|
# - userdel -r <name> 删除账号及家目录
|
||||||
|
# - passwd -l / -u <name> 锁定/解锁口令
|
||||||
|
# - 不授予 chsh/其他命令的任意执行;若需变更默认 shell 请收紧为固定参数
|
||||||
|
#
|
||||||
|
# 生产禁止 system.sudo=false 的 direct 模式:必须显式配置
|
||||||
|
# [system]
|
||||||
|
# sudo = true
|
||||||
|
# dry_run = false
|
||||||
@@ -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)")
|
|
||||||
}
|
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
+41
-10
@@ -1,36 +1,57 @@
|
|||||||
// Package api 为 HTTP handler 层(RESTful v1)。
|
// Package api 为 HTTP handler 层(RESTful v1)。
|
||||||
// M0 提供健康检查与模块路由骨架;各模块 handler 在对应里程碑填充。
|
// M1 覆盖认证(管理员/外部用户登录、会话、OTP)与用户管理 CRUD。
|
||||||
package api
|
package api
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"ws_usernode/internal/auth"
|
||||||
|
"ws_usernode/internal/config"
|
||||||
"ws_usernode/internal/service"
|
"ws_usernode/internal/service"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Handler 聚合各模块 handler,作为路由注册的挂载点。
|
// Handler 聚合各模块 handler,作为路由注册的挂载点。
|
||||||
type Handler struct {
|
type Handler struct {
|
||||||
Health *HealthHandler
|
Health *HealthHandler
|
||||||
Admin *AdminHandler
|
Auth *AuthHandler
|
||||||
User *UserHandler
|
User *UserHandler
|
||||||
// Auth / Keys / Approval / Audit / Settings 等模块在 M1~M4 填充
|
|
||||||
|
authSvc *service.AuthService
|
||||||
|
auditSvc *service.AuditService
|
||||||
}
|
}
|
||||||
|
|
||||||
// New 创建 handler 集合。M0 阶段部分服务可为 nil,路由只挂已实现模块。
|
// New 创建 handler 集合。
|
||||||
func New(adminSvc *service.AdminService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler {
|
func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler {
|
||||||
h := &Handler{
|
h := &Handler{
|
||||||
Health: &HealthHandler{startedAt: time.Now()},
|
Health: &HealthHandler{startedAt: time.Now()},
|
||||||
Admin: &AdminHandler{svc: adminSvc},
|
authSvc: authSvc,
|
||||||
User: &UserHandler{svc: userSvc},
|
auditSvc: auditSvc,
|
||||||
}
|
}
|
||||||
_ = auditSvc
|
h.Auth = &AuthHandler{svc: authSvc, cfg: cfg}
|
||||||
|
h.User = &UserHandler{svc: userSvc, cfg: cfg, h: h}
|
||||||
return 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 健康检查。
|
// HealthHandler 健康检查。
|
||||||
type HealthHandler struct {
|
type HealthHandler struct {
|
||||||
startedAt time.Time
|
startedAt time.Time
|
||||||
@@ -39,13 +60,23 @@ type HealthHandler struct {
|
|||||||
func (h *HealthHandler) Healthz(c *gin.Context) {
|
func (h *HealthHandler) Healthz(c *gin.Context) {
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"status": "ok",
|
"status": "ok",
|
||||||
"version": "0.1.0-m0",
|
"version": "0.2.0-m1",
|
||||||
"uptime": time.Since(h.startedAt).String(),
|
"uptime": time.Since(h.startedAt).String(),
|
||||||
"go": runtime.Version(),
|
"go": runtime.Version(),
|
||||||
"timestamp": time.Now().UTC().Format(time.RFC3339),
|
"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 统一成功响应。
|
// ok 统一成功响应。
|
||||||
func ok(c *gin.Context, data any) {
|
func ok(c *gin.Context, data any) {
|
||||||
c.JSON(http.StatusOK, gin.H{"data": data})
|
c.JSON(http.StatusOK, gin.H{"data": data})
|
||||||
|
|||||||
+176
-13
@@ -1,17 +1,24 @@
|
|||||||
package api
|
package api
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"ws_usernode/internal/config"
|
||||||
|
"ws_usernode/internal/model"
|
||||||
"ws_usernode/internal/service"
|
"ws_usernode/internal/service"
|
||||||
)
|
)
|
||||||
|
|
||||||
// UserHandler 外部用户接口(列表/详情/创建等,M1 填充 CRUD 与系统操作)。
|
// UserHandler 外部用户接口(列表/详情/创建/更新/禁用/启用/延期/删除,admin)。
|
||||||
type UserHandler struct {
|
type UserHandler struct {
|
||||||
svc *service.UserService
|
svc *service.UserService
|
||||||
|
cfg *config.Config
|
||||||
|
h *Handler // 访问审计 helper
|
||||||
}
|
}
|
||||||
|
|
||||||
// UserCreateRequest 管理员创建外部用户请求。
|
// UserCreateRequest 管理员创建外部用户请求。
|
||||||
@@ -25,31 +32,187 @@ type UserCreateRequest struct {
|
|||||||
|
|
||||||
// Create 管理员创建外部用户(自动建系统账号)。
|
// Create 管理员创建外部用户(自动建系统账号)。
|
||||||
func (h *UserHandler) Create(c *gin.Context) {
|
func (h *UserHandler) Create(c *gin.Context) {
|
||||||
if h.svc == nil {
|
|
||||||
fail(c, http.StatusNotImplemented, "用户服务未初始化")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
var req UserCreateRequest
|
var req UserCreateRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error())
|
fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error())
|
||||||
return
|
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 {
|
if err != nil {
|
||||||
|
h.h.audit(c, "user.create", "user", "", map[string]any{"username": req.Username, "err": err.Error()}, model.ResultFailed)
|
||||||
switch {
|
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())
|
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())
|
fail(c, http.StatusConflict, err.Error())
|
||||||
default:
|
default:
|
||||||
fail(c, http.StatusInternalServerError, err.Error())
|
fail(c, http.StatusInternalServerError, err.Error())
|
||||||
}
|
}
|
||||||
return
|
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 实现分页筛选)。
|
// ExtendRequest 延期请求。
|
||||||
func (h *UserHandler) List(c *gin.Context) {
|
type ExtendRequest struct {
|
||||||
fail(c, http.StatusNotImplemented, "用户列表将在 M1 实现")
|
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"})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,10 +14,11 @@ var ErrCaptchaInvalid = errors.New("auth: 图形验证码错误")
|
|||||||
// Captcha 图形验证码(防机器人,登录前置)。
|
// Captcha 图形验证码(防机器人,登录前置)。
|
||||||
type Captcha struct {
|
type Captcha struct {
|
||||||
ID string
|
ID string
|
||||||
Text string // M1 生成图像渲染,此处仅存文本
|
Text string
|
||||||
}
|
}
|
||||||
|
|
||||||
// CaptchaStore 为图形验证码存储(M1 实现图像渲染)。
|
// CaptchaStore 为图形验证码存储。单实例内存实现为默认;
|
||||||
|
// 多实例部署需改 DB/Redis(PLAN §6 注明)。
|
||||||
type CaptchaStore interface {
|
type CaptchaStore interface {
|
||||||
// New 生成一个验证码并返回其 ID。
|
// New 生成一个验证码并返回其 ID。
|
||||||
New() (*Captcha, error)
|
New() (*Captcha, error)
|
||||||
@@ -27,6 +28,7 @@ type CaptchaStore interface {
|
|||||||
|
|
||||||
// MemoryCaptchaStore 单实例内存实现。
|
// MemoryCaptchaStore 单实例内存实现。
|
||||||
type MemoryCaptchaStore struct {
|
type MemoryCaptchaStore struct {
|
||||||
|
ttl time.Duration
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
entries map[string]*captchaEntry
|
entries map[string]*captchaEntry
|
||||||
}
|
}
|
||||||
@@ -37,8 +39,8 @@ type captchaEntry struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NewMemoryCaptchaStore 创建内存图形验证码存储。
|
// NewMemoryCaptchaStore 创建内存图形验证码存储。
|
||||||
func NewMemoryCaptchaStore() *MemoryCaptchaStore {
|
func NewMemoryCaptchaStore(ttl time.Duration) *MemoryCaptchaStore {
|
||||||
return &MemoryCaptchaStore{entries: make(map[string]*captchaEntry)}
|
return &MemoryCaptchaStore{ttl: ttl, entries: make(map[string]*captchaEntry)}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *MemoryCaptchaStore) New() (*Captcha, error) {
|
func (s *MemoryCaptchaStore) New() (*Captcha, error) {
|
||||||
@@ -52,7 +54,7 @@ func (s *MemoryCaptchaStore) New() (*Captcha, error) {
|
|||||||
}
|
}
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
defer s.mu.Unlock()
|
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
|
return &Captcha{ID: id, Text: text}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+159
-17
@@ -1,15 +1,22 @@
|
|||||||
// Package auth 提供认证相关能力:OTP(双通道)、会话、bcrypt、图形验证码。
|
// Package auth 提供认证相关能力:OTP(双通道)、会话、bcrypt、图形验证码。
|
||||||
//
|
//
|
||||||
// OTP 双通道对齐:邮件发送与 CLI 获取共用同一 OTPStore(同一验证码、同一
|
// OTP 双通道对齐:邮件发送与 CLI 获取共用同一 OTPStore(同一验证码、同一
|
||||||
// 10 分钟有效期、同一 60s 冷却与失败限速),邮件失败不阻断 CLI 通道。
|
// 有效期、同一冷却与失败限速),邮件失败不阻断 CLI 通道。邮件通道经
|
||||||
// 单实例用内存存储;多实例需改为 DB/Redis(PLAN §6 注明)。
|
// Send 生成验证码,CLI 通道优先经 Current 复用同一验证码,无有效码时才
|
||||||
|
// 触发 Send(仍受同一冷却约束)。验证码落 DB(otp_codes 表),单实例部署
|
||||||
|
// 即可保证跨进程(HTTP 服务与 CLI 子命令)共享;多实例需改 DB 行锁/Redis
|
||||||
|
// (PLAN §6)。
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"ws_usernode/internal/model"
|
||||||
"ws_usernode/internal/pkg"
|
"ws_usernode/internal/pkg"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -21,20 +28,141 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
otpCodeLen = 6
|
otpCodeLen = 6
|
||||||
maxFailures = 5 // 单账号连续失败限速阈值
|
|
||||||
failureWin = 10 * time.Minute // 失败计数窗口
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// OTPStore 为 OTP 验证码存储。内存实现为单实例默认实现。
|
// OTP 失败限速默认参数(单账号连续失败阈值与计数窗口)。
|
||||||
|
const (
|
||||||
|
DefaultMaxFailures = 5
|
||||||
|
DefaultFailureWin = 10 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
// OTPStore 为 OTP 验证码存储。DB 实现(DBOTPStore)为生产默认,
|
||||||
|
// MemoryOTPStore 供测试与单进程内嵌场景使用。
|
||||||
type OTPStore interface {
|
type OTPStore interface {
|
||||||
// Send 为 username 生成新验证码(覆盖旧码)。冷却期内调用返回 ErrCooldown。
|
// Send 为 username 生成新验证码并覆盖旧码(邮件通道)。冷却期内返回 ErrCooldown。
|
||||||
// 邮件与 CLI 双通道都走该方法,保证对齐。
|
Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error)
|
||||||
Send(username string, ttl, cooldown time.Duration) (string, error)
|
// Current 返回当前有效(未过期、未消费)验证码,供 CLI 通道复用同一验证码。
|
||||||
|
// 无有效验证码返回 ErrInvalidCode。
|
||||||
|
Current(ctx context.Context, username string) (string, error)
|
||||||
// Verify 校验验证码并一次性消费。失败累计计数(达到阈值返回 ErrTooManyFails)。
|
// Verify 校验验证码并一次性消费。失败累计计数(达到阈值返回 ErrTooManyFails)。
|
||||||
Verify(username, code string) (bool, error)
|
Verify(ctx context.Context, username, code string) (bool, error)
|
||||||
// Failures 返回 username 当前失败计数。
|
// 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 {
|
type otpEntry struct {
|
||||||
@@ -44,7 +172,7 @@ type otpEntry struct {
|
|||||||
failures int
|
failures int
|
||||||
}
|
}
|
||||||
|
|
||||||
// MemoryOTPStore 为单实例内存实现。
|
// MemoryOTPStore 为单进程内存实现(测试/内嵌场景)。
|
||||||
type MemoryOTPStore struct {
|
type MemoryOTPStore struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
entries map[string]*otpEntry
|
entries map[string]*otpEntry
|
||||||
@@ -55,9 +183,10 @@ func NewMemoryOTPStore() *MemoryOTPStore {
|
|||||||
return &MemoryOTPStore{entries: make(map[string]*otpEntry)}
|
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()
|
s.mu.Lock()
|
||||||
defer s.mu.Unlock()
|
defer s.mu.Unlock()
|
||||||
|
_ = ctx
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
if e, ok := s.entries[username]; ok && now.Before(e.cooldownAt) {
|
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
|
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()
|
s.mu.Lock()
|
||||||
defer s.mu.Unlock()
|
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]
|
e, ok := s.entries[username]
|
||||||
if !ok {
|
if !ok {
|
||||||
@@ -89,12 +230,12 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
|
|||||||
delete(s.entries, username)
|
delete(s.entries, username)
|
||||||
return false, ErrInvalidCode
|
return false, ErrInvalidCode
|
||||||
}
|
}
|
||||||
if e.failures >= maxFailures {
|
if e.failures >= DefaultMaxFailures {
|
||||||
return false, ErrTooManyFails
|
return false, ErrTooManyFails
|
||||||
}
|
}
|
||||||
if e.code != code {
|
if e.code != code {
|
||||||
e.failures++
|
e.failures++
|
||||||
if e.failures >= maxFailures {
|
if e.failures >= DefaultMaxFailures {
|
||||||
return false, ErrTooManyFails
|
return false, ErrTooManyFails
|
||||||
}
|
}
|
||||||
return false, ErrInvalidCode
|
return false, ErrInvalidCode
|
||||||
@@ -103,9 +244,10 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
|
|||||||
return true, nil
|
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()
|
s.mu.Lock()
|
||||||
defer s.mu.Unlock()
|
defer s.mu.Unlock()
|
||||||
|
_ = ctx
|
||||||
if e, ok := s.entries[username]; ok {
|
if e, ok := s.entries[username]; ok {
|
||||||
return e.failures, nil
|
return e.failures, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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[:])
|
||||||
|
}
|
||||||
+29
-11
@@ -28,13 +28,15 @@ type Config struct {
|
|||||||
Database DatabaseConfig `toml:"database"`
|
Database DatabaseConfig `toml:"database"`
|
||||||
Log LogConfig `toml:"log"`
|
Log LogConfig `toml:"log"`
|
||||||
Policy PolicyConfig `toml:"policy"`
|
Policy PolicyConfig `toml:"policy"`
|
||||||
|
Auth AuthConfig `toml:"auth"`
|
||||||
SMTP SMTPConfig `toml:"smtp"`
|
SMTP SMTPConfig `toml:"smtp"`
|
||||||
System SystemConfig `toml:"system"`
|
System SystemConfig `toml:"system"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type AppConfig struct {
|
type AppConfig struct {
|
||||||
Name string `toml:"name"`
|
Name string `toml:"name"`
|
||||||
Env string `toml:"env"` // development / production
|
Env string `toml:"env"` // development / production
|
||||||
|
BaseURL string `toml:"base_url"` // 对外访问地址(邮件中的链接使用)
|
||||||
}
|
}
|
||||||
|
|
||||||
type ServerConfig struct {
|
type ServerConfig struct {
|
||||||
@@ -62,6 +64,13 @@ type PolicyConfig struct {
|
|||||||
OTPCooldown time.Duration `toml:"otp_cooldown"` // OTP 发送冷却
|
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 {
|
type SMTPConfig struct {
|
||||||
Host string `toml:"host"`
|
Host string `toml:"host"`
|
||||||
Port int `toml:"port"`
|
Port int `toml:"port"`
|
||||||
@@ -72,21 +81,22 @@ type SMTPConfig struct {
|
|||||||
|
|
||||||
// SystemConfig 为系统账号操作层的本地实现配置(sudoers 白名单模式)。
|
// SystemConfig 为系统账号操作层的本地实现配置(sudoers 白名单模式)。
|
||||||
type SystemConfig struct {
|
type SystemConfig struct {
|
||||||
Sudo bool `toml:"sudo"` // 是否通过 sudo -n 执行系统命令;开发环境 false 时 dry-run
|
Sudo bool `toml:"sudo"` // 是否通过 sudo -n 执行系统命令(生产)
|
||||||
UserPrefix string `toml:"user_prefix"` // 外部用户系统账号前缀,默认 ext_
|
DryRun bool `toml:"dry_run"` // true = 只打印计划命令不执行(开发演练);false 且 sudo=false 时直接执行(容器/测试用户验证)
|
||||||
Group string `toml:"group"` // 外部用户所属组,默认 external
|
UserPrefix string `toml:"user_prefix"` // 外部用户系统账号前缀,默认 ext_
|
||||||
Shell string `toml:"shell"` // 默认 shell
|
Group string `toml:"group"` // 外部用户所属组,默认 external
|
||||||
HomeBase string `toml:"home_base"` // 家目录基路径
|
Shell string `toml:"shell"` // 默认 shell
|
||||||
AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖)
|
HomeBase string `toml:"home_base"` // 家目录基路径
|
||||||
|
AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Default 返回带开发环境默认值的配置,作为 config.example.toml 与未配置项的兜底。
|
// Default 返回带开发环境默认值的配置,作为 config.example.toml 与未配置项的兜底。
|
||||||
func Default() *Config {
|
func Default() *Config {
|
||||||
return &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{
|
Server: ServerConfig{
|
||||||
Listen: "127.0.0.1:8080",
|
Listen: "127.0.0.1:8080",
|
||||||
SessionTTL: 24 * time.Hour,
|
SessionTTL: 24 * time.Hour,
|
||||||
TrustedProxies: []string{"127.0.0.1", "::1"},
|
TrustedProxies: []string{"127.0.0.1", "::1"},
|
||||||
},
|
},
|
||||||
Database: DatabaseConfig{Driver: "sqlite", DSN: "data/usernode.db"},
|
Database: DatabaseConfig{Driver: "sqlite", DSN: "data/usernode.db"},
|
||||||
@@ -98,9 +108,17 @@ func Default() *Config {
|
|||||||
OTPTTL: 10 * time.Minute,
|
OTPTTL: 10 * time.Minute,
|
||||||
OTPCooldown: 60 * time.Second,
|
OTPCooldown: 60 * time.Second,
|
||||||
},
|
},
|
||||||
|
Auth: AuthConfig{
|
||||||
|
MaxLoginFailures: 5,
|
||||||
|
LockDuration: 15 * time.Minute,
|
||||||
|
CaptchaTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
SMTP: SMTPConfig{Port: 587},
|
SMTP: SMTPConfig{Port: 587},
|
||||||
System: SystemConfig{
|
System: SystemConfig{
|
||||||
|
// 开发默认 dry-run:未配置 config 直接跑 serve 时只打印计划,避免误操作系统账号。
|
||||||
|
// 生产必须显式 dry_run=false 且 sudo=true(见 deploy/sudoers.example)。
|
||||||
Sudo: false,
|
Sudo: false,
|
||||||
|
DryRun: true,
|
||||||
UserPrefix: "ext_",
|
UserPrefix: "ext_",
|
||||||
Group: "external",
|
Group: "external",
|
||||||
Shell: "/bin/sh",
|
Shell: "/bin/sh",
|
||||||
|
|||||||
@@ -31,6 +31,12 @@ func TestLoadDefault(t *testing.T) {
|
|||||||
if cfg.System.UserPrefix != "ext_" {
|
if cfg.System.UserPrefix != "ext_" {
|
||||||
t.Errorf("default user_prefix = %q", cfg.System.UserPrefix)
|
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) {
|
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_DATABASE_DSN", "u:p@tcp(h:3306)/db")
|
||||||
t.Setenv("USERNODE_POLICY_OTPTTL", "5m")
|
t.Setenv("USERNODE_POLICY_OTPTTL", "5m")
|
||||||
t.Setenv("USERNODE_SYSTEM_SUDO", "true")
|
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_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()
|
cfg, err := LoadDefault()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -77,6 +87,12 @@ func TestEnvOverrides(t *testing.T) {
|
|||||||
if !cfg.System.Sudo {
|
if !cfg.System.Sudo {
|
||||||
t.Error("system.sudo should be true")
|
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" {
|
if len(cfg.Server.TrustedProxies) != 2 || cfg.Server.TrustedProxies[0] != "10.0.0.1" {
|
||||||
t.Errorf("trusted_proxies = %v", cfg.Server.TrustedProxies)
|
t.Errorf("trusted_proxies = %v", cfg.Server.TrustedProxies)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -141,11 +141,43 @@ type MailLog struct {
|
|||||||
UpdatedAt time.Time `json:"updated_at"`
|
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 使用的全部模型。
|
// AllModels 供 AutoMigrate 使用的全部模型。
|
||||||
func AllModels() []any {
|
func AllModels() []any {
|
||||||
return []any{
|
return []any{
|
||||||
&AdminUser{}, &User{}, &SSHKey{}, &Approval{},
|
&AdminUser{}, &User{}, &SSHKey{}, &Approval{},
|
||||||
&AuditLog{}, &Session{}, &Setting{}, &MailLog{},
|
&AuditLog{}, &Session{}, &Setting{}, &MailLog{},
|
||||||
|
&OTPCode{}, &PasswordResetToken{},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
|
}
|
||||||
+28
-14
@@ -9,13 +9,14 @@ import (
|
|||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"ws_usernode/internal/api"
|
"ws_usernode/internal/api"
|
||||||
|
"ws_usernode/internal/auth"
|
||||||
"ws_usernode/internal/config"
|
"ws_usernode/internal/config"
|
||||||
"ws_usernode/internal/webui"
|
"ws_usernode/internal/webui"
|
||||||
)
|
)
|
||||||
|
|
||||||
// New 构建根 router:API v1 + 前端静态资源(go:embed)。
|
// New 构建根 router:API v1 + 前端静态资源(go:embed)。
|
||||||
// production 模式启用 gin.ReleaseMode;否则启用调试模式与开发日志。
|
// 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" {
|
if cfg.App.Env == "production" {
|
||||||
gin.SetMode(gin.ReleaseMode)
|
gin.SetMode(gin.ReleaseMode)
|
||||||
}
|
}
|
||||||
@@ -28,21 +29,34 @@ func New(cfg *config.Config, h *api.Handler, log *slog.Logger) *gin.Engine {
|
|||||||
// RESTful API v1
|
// RESTful API v1
|
||||||
v1 := r.Group("/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
|
authGrp.GET("/captcha", h.Auth.Captcha)
|
||||||
auth.GET("/captcha", notImplemented("图形验证码(M1)"))
|
authGrp.POST("/otp/send", h.Auth.OTPSend)
|
||||||
auth.POST("/otp/send", notImplemented("OTP 发送(M1)"))
|
authGrp.POST("/otp/login", h.Auth.OTPLogin)
|
||||||
auth.POST("/otp/login", notImplemented("OTP 登录(M1)"))
|
authGrp.POST("/admin/login", h.Auth.AdminLogin)
|
||||||
auth.POST("/admin/login", notImplemented("管理员登录(M1)"))
|
authGrp.POST("/admin/forgot", h.Auth.AdminForgot)
|
||||||
auth.POST("/logout", notImplemented("登出(M1)"))
|
authGrp.POST("/admin/reset", h.Auth.AdminReset)
|
||||||
auth.GET("/me", notImplemented("当前会话(M1)"))
|
// 需要会话(管理员或外部用户)
|
||||||
|
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)
|
// 用户管理(admin)
|
||||||
v1.POST("/users/:id/disable", notImplemented("禁用用户(M1)"))
|
users := v1.Group("/users", sessionMiddleware(sessions), requireUserType(auth.SessionUserAdmin))
|
||||||
v1.POST("/users/:id/enable", notImplemented("启用用户(M1)"))
|
{
|
||||||
v1.POST("/users/:id/extend", notImplemented("延期(M1)"))
|
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.POST("/approvals", notImplemented("提交申请(M3)"))
|
||||||
v1.GET("/approvals", notImplemented("申请列表(M3)"))
|
v1.GET("/approvals", notImplemented("申请列表(M3)"))
|
||||||
v1.POST("/approvals/:id/review", notImplemented("审批(M3)"))
|
v1.POST("/approvals/:id/review", notImplemented("审批(M3)"))
|
||||||
|
|||||||
@@ -15,9 +15,10 @@ import (
|
|||||||
|
|
||||||
// 管理员服务错误。
|
// 管理员服务错误。
|
||||||
var (
|
var (
|
||||||
ErrAdminExists = errors.New("service: 管理员已存在")
|
ErrAdminExists = errors.New("service: 管理员已存在")
|
||||||
ErrAdminNotFound = errors.New("service: 管理员不存在")
|
ErrAdminNotFound = errors.New("service: 管理员不存在")
|
||||||
ErrWeakPassword = errors.New("service: 密码过弱(至少 8 位,需含字母与数字)")
|
ErrWeakPassword = errors.New("service: 密码过弱(至少 8 位,需含字母与数字)")
|
||||||
|
ErrBadCredentials = errors.New("service: 用户名或密码错误")
|
||||||
)
|
)
|
||||||
|
|
||||||
// AdminService 管理端账号服务(CLI admin create / reset-password 与登录共用)。
|
// AdminService 管理端账号服务(CLI admin create / reset-password 与登录共用)。
|
||||||
@@ -94,6 +95,43 @@ func (s *AdminService) GetByUsername(ctx context.Context, username string) (*mod
|
|||||||
return &adm, nil
|
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 可加策略)。
|
// validatePassword 校验管理员密码强度(骨架阶段基础规则,M1 可加策略)。
|
||||||
func validatePassword(p string) error {
|
func validatePassword(p string) error {
|
||||||
if len(p) < 8 {
|
if len(p) < 8 {
|
||||||
|
|||||||
@@ -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: 会话无效")
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+229
-17
@@ -4,9 +4,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"ws_usernode/internal/config"
|
||||||
"ws_usernode/internal/model"
|
"ws_usernode/internal/model"
|
||||||
"ws_usernode/internal/pkg"
|
"ws_usernode/internal/pkg"
|
||||||
"ws_usernode/internal/system"
|
"ws_usernode/internal/system"
|
||||||
@@ -14,29 +16,30 @@ import (
|
|||||||
|
|
||||||
// 用户服务错误。
|
// 用户服务错误。
|
||||||
var (
|
var (
|
||||||
ErrUserNotFound = errors.New("service: 用户不存在")
|
ErrUserNotFound = errors.New("service: 用户不存在")
|
||||||
ErrUserExists = errors.New("service: 用户名已存在")
|
ErrUserExists = errors.New("service: 用户名已存在")
|
||||||
|
ErrUserExpired = errors.New("service: 用户已过期,请先延期")
|
||||||
|
ErrUserDisabled = errors.New("service: 用户已禁用,无法操作")
|
||||||
|
ErrSystemAccountMissing = errors.New("service: 系统账号不存在,无法操作")
|
||||||
)
|
)
|
||||||
|
|
||||||
// UserService 外部用户生命周期服务。
|
// UserService 外部用户生命周期服务:DB 记录 + system.Manager 系统账号操作。
|
||||||
// M0 提供查询与创建骨架;建号(useradd)、禁用/延期等系统操作 M1 接入 system.Manager。
|
|
||||||
type UserService struct {
|
type UserService struct {
|
||||||
db *gorm.DB
|
db *gorm.DB
|
||||||
sys system.Manager
|
sys system.Manager
|
||||||
|
cfg *config.Config
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewUserService 创建用户服务。
|
// NewUserService 创建用户服务。
|
||||||
func NewUserService(db *gorm.DB, sys system.Manager) *UserService {
|
func NewUserService(db *gorm.DB, sys system.Manager, cfg *config.Config) *UserService {
|
||||||
return &UserService{db: db, sys: sys}
|
return &UserService{db: db, sys: sys, cfg: cfg}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetByUsername 按用户名查询外部用户(含或不含 ext_ 前缀均可)。
|
// GetByUsername 按用户名查询外部用户(含或不含 ext_ 前缀均可)。
|
||||||
func (s *UserService) GetByUsername(ctx context.Context, username string) (*model.User, error) {
|
func (s *UserService) GetByUsername(ctx context.Context, username string) (*model.User, error) {
|
||||||
if !strings.HasPrefix(username, "ext_") {
|
full := normalizeName(username, s.cfg.System.UserPrefix)
|
||||||
username = "ext_" + username
|
|
||||||
}
|
|
||||||
var u model.User
|
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) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return nil, ErrUserNotFound
|
return nil, ErrUserNotFound
|
||||||
}
|
}
|
||||||
@@ -45,16 +48,68 @@ func (s *UserService) GetByUsername(ctx context.Context, username string) (*mode
|
|||||||
return &u, nil
|
return &u, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create 创建外部用户记录并调用系统层建号。M0 阶段系统层为 dry-run。
|
// GetByID 按 ID 查询外部用户。
|
||||||
// username 为不含前缀的申请名,内部加 ext_ 前缀。
|
func (s *UserService) GetByID(ctx context.Context, id uint) (*model.User, error) {
|
||||||
func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttlSeconds int64) (*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 {
|
if err := pkg.ValidateUserName(username); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := pkg.ValidateEmail(email); err != nil {
|
if err := pkg.ValidateEmail(email); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
full := "ext_" + username
|
username = strings.TrimSpace(username)
|
||||||
|
full := s.cfg.System.UserPrefix + username
|
||||||
var count int64
|
var count int64
|
||||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", full).Count(&count).Error; err != nil {
|
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", full).Count(&count).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -62,21 +117,178 @@ func (s *UserService) Create(ctx context.Context, username, email, supervisor, p
|
|||||||
if count > 0 {
|
if count > 0 {
|
||||||
return nil, ErrUserExists
|
return nil, ErrUserExists
|
||||||
}
|
}
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = s.cfg.Policy.DefaultTTL
|
||||||
|
}
|
||||||
|
expireAt := time.Now().Add(ttl)
|
||||||
u := &model.User{
|
u := &model.User{
|
||||||
Username: full,
|
Username: full,
|
||||||
Email: email,
|
Email: email,
|
||||||
Supervisor: supervisor,
|
Supervisor: supervisor,
|
||||||
Purpose: purpose,
|
Purpose: purpose,
|
||||||
Status: model.UserStatusActive,
|
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 {
|
if err := s.db.WithContext(ctx).Create(u).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
// 系统账号创建(dry-run / 真实),失败时回滚 DB 记录
|
// 系统账号创建(dry-run / 直接 / sudo),失败时回滚 DB 记录
|
||||||
if err := s.sys.CreateUser(ctx, system.Account{Username: full}); err != nil {
|
if err := s.sys.CreateUser(ctx, system.Account{Username: full, Shell: s.cfg.System.Shell}); err != nil {
|
||||||
_ = s.db.WithContext(ctx).Delete(u).Error
|
_ = s.db.WithContext(ctx).Delete(u).Error
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return u, nil
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+71
-22
@@ -4,15 +4,19 @@
|
|||||||
// 执行 useradd/usermod/userdel/passwd 等固定命令并做参数强校验。未来多节点
|
// 执行 useradd/usermod/userdel/passwd 等固定命令并做参数强校验。未来多节点
|
||||||
// agent 模式只需新增远程实现替换本地实现(PLAN §5.1)。
|
// agent 模式只需新增远程实现替换本地实现(PLAN §5.1)。
|
||||||
//
|
//
|
||||||
// 权限模型:节点以专有用户(如 usernode)运行,经 sudo -n 提权执行白名单
|
// 执行模式(config system.*):
|
||||||
// 命令;开发环境 config system.sudo=false 进入 dry-run(只打印计划不执行),
|
// - dry_run=true:只打印计划命令不执行(开发演练);
|
||||||
// 避免在开发机上直接操作系统账号。系统命令的真实系统效果在 M1 用测试用户/
|
// - dry_run=false + sudo=false:直接执行(容器/测试用户验证真实建号);
|
||||||
// 容器验证,不在生产直接跑 useradd。
|
// - dry_run=false + sudo=true:经 sudo -n 提权执行白名单命令(生产,
|
||||||
|
// 需配置 deploy/sudoers)。真实系统账号操作只在测试用户/容器中验证,
|
||||||
|
// 不在生产直接跑 useradd。
|
||||||
package system
|
package system
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -45,8 +49,10 @@ type Manager interface {
|
|||||||
RemoveUser(ctx context.Context, username string) error
|
RemoveUser(ctx context.Context, username string) error
|
||||||
// SetLock 锁定/解锁系统账号口令(passwd -l / -u)。
|
// SetLock 锁定/解锁系统账号口令(passwd -l / -u)。
|
||||||
SetLock(ctx context.Context, username string, locked bool) error
|
SetLock(ctx context.Context, username string, locked bool) error
|
||||||
|
// Exists 检查系统账号是否存在(读 /etc/passwd,无需提权)。
|
||||||
|
Exists(ctx context.Context, username string) (bool, error)
|
||||||
// SyncAuthorizedKeys 以 DB 状态全量重写 authorized_keys(原子写 + 并发锁),
|
// SyncAuthorizedKeys 以 DB 状态全量重写 authorized_keys(原子写 + 并发锁),
|
||||||
// 吊销密钥即从文件移除、立即失效。M0 提供 dry-run 实现,M2 完成生产路径。
|
// 吊销密钥即从文件移除、立即失效。
|
||||||
SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error
|
SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -69,20 +75,32 @@ func New(cfg config.SystemConfig) Manager {
|
|||||||
return &localManager{cfg: cfg}
|
return &localManager{cfg: cfg}
|
||||||
}
|
}
|
||||||
|
|
||||||
// run 执行白名单命令:sudo -n <cmd> <args...>,参数在调用处强校验。
|
// run 执行白名单命令。参数在调用处强校验。
|
||||||
// dry-run 模式下只返回将执行的命令文本,不真正执行。
|
// 模式:dry-run 只返回命令文本;direct(sudo=false)直接执行;
|
||||||
|
// sudo 经 sudo -n 提权执行(对应 deploy/sudoers 白名单)。
|
||||||
func (m *localManager) run(ctx context.Context, cmd string, args ...string) (string, error) {
|
func (m *localManager) run(ctx context.Context, cmd string, args ...string) (string, error) {
|
||||||
if !allowedCommands[cmd] {
|
if !allowedCommands[cmd] {
|
||||||
return "", errors.New("system: command not allowed: " + cmd)
|
return "", errors.New("system: command not allowed: " + cmd)
|
||||||
}
|
}
|
||||||
argv := append([]string{"-n", cmd}, args...)
|
cmdline := strings.Join(append([]string{cmd}, args...), " ")
|
||||||
cmdline := strings.Join(append([]string{"sudo", "-n", cmd}, args...), " ")
|
if m.cfg.DryRun {
|
||||||
if !m.cfg.Sudo {
|
return cmdline, nil
|
||||||
return cmdline, nil // dry-run
|
}
|
||||||
|
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 {
|
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
|
return cmdline, nil
|
||||||
}
|
}
|
||||||
@@ -109,11 +127,11 @@ func (m *localManager) CreateUser(ctx context.Context, acc Account) error {
|
|||||||
home = filepath.Join(m.cfg.HomeBase, username)
|
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 {
|
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 {
|
if _, err := m.run(ctx, "passwd", "-l", username); err != nil {
|
||||||
return errors.New("system: passwd -l: " + err.Error())
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -124,7 +142,7 @@ func (m *localManager) RemoveUser(ctx context.Context, username string) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if _, err := m.run(ctx, "userdel", "-r", username); err != nil {
|
if _, err := m.run(ctx, "userdel", "-r", username); err != nil {
|
||||||
return errors.New("system: userdel: " + err.Error())
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -139,20 +157,51 @@ func (m *localManager) SetLock(ctx context.Context, username string, locked bool
|
|||||||
flag = "-l"
|
flag = "-l"
|
||||||
}
|
}
|
||||||
if _, err := m.run(ctx, "passwd", flag, username); err != nil {
|
if _, err := m.run(ctx, "passwd", flag, username); err != nil {
|
||||||
return errors.New("system: passwd: " + err.Error())
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
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:
|
// SyncAuthorizedKeys 全量重写 authorized_keys:
|
||||||
//
|
//
|
||||||
// 1. 以 <lockBase>/<username>.lock 文件锁串行化并发写(flock);
|
// 1. 以 <lockBase>/<username>.lock 文件锁串行化并发写(flock);
|
||||||
// 2. 写临时文件(0600),再 rename 原子替换;
|
// 2. 写临时文件(0600),再 rename 原子替换;
|
||||||
// 3. 无有效密钥时写入空文件(SSH 行为一致,避免文件缺失导致的歧义)。
|
// 3. 无有效密钥时写入空文件(SSH 行为一致,避免文件缺失导致的歧义)。
|
||||||
//
|
//
|
||||||
// 生产模式(sudo=true)下节点进程需具备对家目录 .ssh 的写权限——部署时经
|
// dry-run 模式只打印计划;真实写路径(direct/sudo)需要进程具备对家目录
|
||||||
// sudoers 白名单授予固定命令/受控脚本实现(M2 细化);当前实现直接做文件
|
// .ssh 的写权限——生产部署经 sudoers 白名单授予受控脚本实现(M2 细化)。
|
||||||
// 操作并假设权限已配置,dry-run 模式打印计划命令。
|
|
||||||
func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error {
|
func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error {
|
||||||
username = m.sysName(username)
|
username = m.sysName(username)
|
||||||
if err := pkg.ValidateSystemAccount(username); err != nil {
|
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
|
var plan strings.Builder
|
||||||
plan.WriteString("mkdir -p " + sshDir + " (0700)\n")
|
plan.WriteString("mkdir -p " + sshDir + " (0700)\n")
|
||||||
plan.WriteString("flock " + lockPath + "\n")
|
plan.WriteString("flock " + lockPath + "\n")
|
||||||
plan.WriteString("write " + filepath.Join(sshDir, "authorized_keys") + " (0600)\n")
|
plan.WriteString("write " + filepath.Join(sshDir, "authorized_keys") + " (0600)\n")
|
||||||
plan.WriteString(content.String())
|
plan.WriteString(content.String())
|
||||||
// 骨架阶段:仅日志输出计划,不落盘
|
// 演练模式:仅日志输出计划,不落盘
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user