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:
2026-08-29 23:40:20 +08:00
parent ae45aba607
commit 630d240dc0
32 changed files with 2923 additions and 188 deletions
+1 -1
View File
@@ -13,7 +13,7 @@ NET_HOST := --network=host
GO ?= go
PODMAN ?= podman
BIN := bin/usernode
VERSION ?= 0.1.0-m0
VERSION ?= 0.2.0-m1
LDFLAGS := -s -w -X main.version=$(VERSION)
GOFLAGS := -trimpath
+14 -3
View File
@@ -10,8 +10,10 @@ import (
"gorm.io/gorm"
"ws_usernode/internal/api"
"ws_usernode/internal/auth"
"ws_usernode/internal/config"
"ws_usernode/internal/cron"
"ws_usernode/internal/mail"
"ws_usernode/internal/model"
"ws_usernode/internal/router"
"ws_usernode/internal/server"
@@ -82,9 +84,18 @@ func cmdServe(args []string) error {
sys := system.New(cfg.System)
adminSvc := service.NewAdminService(db)
userSvc := service.NewUserService(db, sys)
userSvc := service.NewUserService(db, sys, cfg)
auditSvc := service.NewAuditService(db)
// 认证依赖:DB OTP/会话/重置令牌存储 + 内存图形验证码/登录限速器 + 邮件
captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL)
otps := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin)
sessions := auth.NewDBSessionStore(db)
resets := auth.NewDBResetTokenStore(db)
limiter := auth.NewRateLimiter(cfg.Auth.MaxLoginFailures, cfg.Auth.LockDuration)
mailer := mail.New(cfg.SMTP, log)
authSvc := service.NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, log)
// 启动前自动迁移(骨架阶段保证表结构就绪;M5 部署建议显式 migrate)
if err := model.Migrate(db); err != nil {
return fmt.Errorf("数据库迁移: %w", err)
@@ -97,8 +108,8 @@ func cmdServe(args []string) error {
sched.Start()
defer sched.Stop()
h := api.New(adminSvc, userSvc, auditSvc)
r := router.New(cfg, h, log)
h := api.New(cfg, authSvc, userSvc, auditSvc)
r := router.New(cfg, h, sessions, log)
srv := server.New(cfg.Server.Listen, r, log)
if err := srv.Run(); err != nil {
+18 -17
View File
@@ -25,8 +25,8 @@ func cmdUser(args []string) error {
}
}
// userOTP 获取外部用户 OTP 验证码与邮件通道共用同一存储与限速
// (同一验证码、同一 10 分钟有效期、同一 60s 冷却与失败限速)。
// userOTP 获取外部用户 OTP 验证码与邮件通道共用同一 DB 存储与限速
// 已有有效验证码时直接复用(同一验证码),无则生成(受同一 60s 冷却约束)。
func userOTP(args []string) error {
fs := flag.NewFlagSet("user otp", flag.ContinueOnError)
cfgPath, debug := commonFlags(fs)
@@ -48,27 +48,28 @@ func userOTP(args []string) error {
}
// 校验用户存在(不存在时返回友好错误,避免暴露账号是否存在的枚举)
if _, err := service.NewUserService(db, system.New(cfg.System)).GetByUsername(context.Background(), name); err != nil {
if _, err := service.NewUserService(db, system.New(cfg.System), cfg).GetByUsername(context.Background(), name); err != nil {
return fmt.Errorf("用户不存在或不可用: %w", err)
}
// M0 骨架:CLI 独立生成(内存 store 与运行中服务不共享)。
// 生产对齐(同一验证码/冷却/限速跨通道生效)需 OTP 落 DB,M1 实现
// auth.OTPStore 的 DB 实现后,CLI 与邮件通道读写同一存储。
store := auth.NewMemoryOTPStore()
code, err := store.Send(name, cfg.Policy.OTPTTL, cfg.Policy.OTPCooldown)
store := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin)
ctx := context.Background()
code, err := store.Current(ctx, name)
if err != nil {
if err == auth.ErrCooldown {
return fmt.Errorf("发送冷却中,请稍后重试(冷却 %s)", cfg.Policy.OTPCooldown)
// 无有效验证码:生成(先到先得,覆盖旧码;冷却期内返回 ErrCooldown
code, err = store.Send(ctx, name, cfg.Policy.OTPTTL, cfg.Policy.OTPCooldown)
if err != nil {
if err == auth.ErrCooldown {
return fmt.Errorf("发送冷却中,请稍后重试(冷却 %s)", cfg.Policy.OTPCooldown)
}
return err
}
return err
log.Info("otp generated", "username", name, "valid_for", cfg.Policy.OTPTTL.String())
} else {
log.Info("otp reused (与邮件通道同一验证码)", "username", name)
}
log.Info("otp generated",
"username", name,
"valid_for", cfg.Policy.OTPTTL.String(),
"expires_at", time.Now().Add(cfg.Policy.OTPTTL).Format(time.RFC3339),
"hint", "与邮件通道为同一验证码,登录后立即失效",
)
fmt.Printf("OTP for %s: %s\n", name, code)
fmt.Printf("有效期至 %s,登录后立即失效;如需重发请等待冷却 %s 或稍后在网页重新请求。\n",
time.Now().Add(cfg.Policy.OTPTTL).Format(time.RFC3339), cfg.Policy.OTPCooldown)
return nil
}
+8 -1
View File
@@ -8,6 +8,7 @@
[app]
name = "ws_usernode"
env = "development" # development / production
base_url = "http://127.0.0.1:8080" # 对外访问地址(邮件重置链接等)
[server]
listen = "127.0.0.1:8080" # 生产建议 0.0.0.0:8080 并置于反向代理后
@@ -30,6 +31,11 @@ audit_retention = "720h" # 审计保留 30 天,保留前先归档
otp_ttl = "10m" # OTP 验证码有效期
otp_cooldown = "60s" # OTP 发送冷却
[auth]
max_login_failures = 5 # 管理员登录连续失败阈值,达到后锁定
lock_duration = "15m" # 锁定持续时间
captcha_ttl = "5m" # 图形验证码有效期
[smtp]
host = "" # 留空则禁用邮件(OTP 仍可用 CLI 通道获取)
port = 587
@@ -38,7 +44,8 @@ password = ""
from = "usernode@example.com"
[system]
sudo = false # 开发环境 false = dry-run(只打印不执行);生产 true 经 sudo -n 执行
sudo = false # 生产 true经 sudo -n 执行 useradd/usermod/userdel/passwd(需 deploy/sudoers
dry_run = true # 开发演练 true:只打印计划命令不执行;false 且 sudo=false 时直接执行(容器/测试用户验证)
user_prefix = "ext_" # 外部用户系统账号统一前缀
group = "external" # 外部用户统一组
shell = "/bin/sh" # 默认 shell
+2 -1
View File
@@ -49,7 +49,8 @@ RUN CGO_ENABLED=0 go build -trimpath -ldflags="-s -w" -o /out/usernode ./cmd/use
# ---------- 阶段 3:运行镜像 ----------
FROM alpine:3.20
RUN apk add --no-cache ca-certificates tzdata \
# shadow 提供 useradd/usermod/userdel/passwdalpine 默认 busybox 无 useradd
RUN apk add --no-cache ca-certificates tzdata shadow \
&& addgroup -S usernode && adduser -S -G usernode usernode
COPY --from=go-build /out/usernode /usr/local/bin/usernode
+23
View File
@@ -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
-53
View File
@@ -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")
}
+326
View File
@@ -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)
}
+195
View File
@@ -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)
}
+40
View File
@@ -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 写入会话 cookieHttpOnly/SameSite=LaxmaxAge<=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
View File
@@ -1,36 +1,57 @@
// Package api 为 HTTP handler 层(RESTful v1)。
// M0 提供健康检查与模块路由骨架;各模块 handler 在对应里程碑填充
// M1 覆盖认证(管理员/外部用户登录、会话、OTP)与用户管理 CRUD
package api
import (
"net/http"
"runtime"
"strconv"
"time"
"github.com/gin-gonic/gin"
"ws_usernode/internal/auth"
"ws_usernode/internal/config"
"ws_usernode/internal/service"
)
// Handler 聚合各模块 handler,作为路由注册的挂载点。
type Handler struct {
Health *HealthHandler
Admin *AdminHandler
Auth *AuthHandler
User *UserHandler
// Auth / Keys / Approval / Audit / Settings 等模块在 M1~M4 填充
authSvc *service.AuthService
auditSvc *service.AuditService
}
// New 创建 handler 集合。M0 阶段部分服务可为 nil,路由只挂已实现模块。
func New(adminSvc *service.AdminService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler {
// New 创建 handler 集合。
func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler {
h := &Handler{
Health: &HealthHandler{startedAt: time.Now()},
Admin: &AdminHandler{svc: adminSvc},
User: &UserHandler{svc: userSvc},
Health: &HealthHandler{startedAt: time.Now()},
authSvc: authSvc,
auditSvc: auditSvc,
}
_ = auditSvc
h.Auth = &AuthHandler{svc: authSvc, cfg: cfg}
h.User = &UserHandler{svc: userSvc, cfg: cfg, h: h}
return h
}
// audit 记录管理操作审计(append-only)。actor 来自会话中间件。
func (h *Handler) audit(c *gin.Context, action, resourceType, resourceID string, detail any, result string) {
var actorID uint
var actorName string
if sess := sessionFrom(c); sess != nil {
actorID = sess.RefID
if info, err := h.authSvc.Me(c.Request.Context(), sess.ID); err == nil {
actorName = info.Username
} else {
actorName = sess.UserType + "#" + strconv.FormatUint(uint64(sess.RefID), 10)
}
}
_ = h.auditSvc.Record(c.Request.Context(), actorID, actorName, action, resourceType, resourceID, detail, c.ClientIP(), result)
}
// HealthHandler 健康检查。
type HealthHandler struct {
startedAt time.Time
@@ -39,13 +60,23 @@ type HealthHandler struct {
func (h *HealthHandler) Healthz(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"version": "0.1.0-m0",
"version": "0.2.0-m1",
"uptime": time.Since(h.startedAt).String(),
"go": runtime.Version(),
"timestamp": time.Now().UTC().Format(time.RFC3339),
})
}
// sessionFrom 返回会话中间件注入的会话(未登录时为 nil)。
func sessionFrom(c *gin.Context) *auth.Session {
if v, ok := c.Get(SessionContextKey); ok {
if s, ok := v.(*auth.Session); ok {
return s
}
}
return nil
}
// ok 统一成功响应。
func ok(c *gin.Context, data any) {
c.JSON(http.StatusOK, gin.H{"data": data})
+176 -13
View File
@@ -1,17 +1,24 @@
package api
import (
"context"
"errors"
"net/http"
"strings"
"strconv"
"time"
"github.com/gin-gonic/gin"
"ws_usernode/internal/config"
"ws_usernode/internal/model"
"ws_usernode/internal/service"
)
// UserHandler 外部用户接口(列表/详情/创建等,M1 填充 CRUD 与系统操作)。
// UserHandler 外部用户接口(列表/详情/创建/更新/禁用/启用/延期/删除,admin)。
type UserHandler struct {
svc *service.UserService
cfg *config.Config
h *Handler // 访问审计 helper
}
// UserCreateRequest 管理员创建外部用户请求。
@@ -25,31 +32,187 @@ type UserCreateRequest struct {
// Create 管理员创建外部用户(自动建系统账号)。
func (h *UserHandler) Create(c *gin.Context) {
if h.svc == nil {
fail(c, http.StatusNotImplemented, "用户服务未初始化")
return
}
var req UserCreateRequest
if err := c.ShouldBindJSON(&req); err != nil {
fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error())
return
}
u, err := h.svc.Create(c.Request.Context(), req.Username, req.Email, req.Supervisor, req.Purpose, req.TTLDays*86400)
createdBy := uint(0)
if sess := sessionFrom(c); sess != nil {
createdBy = sess.RefID
}
u, err := h.svc.Create(c.Request.Context(), req.Username, req.Email, req.Supervisor, req.Purpose,
time.Duration(req.TTLDays)*24*time.Hour, createdBy)
if err != nil {
h.h.audit(c, "user.create", "user", "", map[string]any{"username": req.Username, "err": err.Error()}, model.ResultFailed)
switch {
case strings.Contains(err.Error(), "用户名"):
case errors.Is(err, service.ErrUserExists):
fail(c, http.StatusConflict, err.Error())
default:
fail(c, http.StatusBadRequest, err.Error())
case strings.Contains(err.Error(), "已存在"):
}
return
}
h.h.audit(c, "user.create", "user", strconv.FormatUint(uint64(u.ID), 10), map[string]any{"username": u.Username}, model.ResultSuccess)
ok(c, gin.H{"id": u.ID, "username": u.Username, "status": u.Status, "expire_at": u.ExpireAt})
}
// List 用户列表(分页/筛选:status、supervisor)。
func (h *UserHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
users, total, err := h.svc.List(c.Request.Context(), service.UserFilter{
Status: c.Query("status"),
Supervisor: c.Query("supervisor"),
Page: page,
PageSize: pageSize,
})
if err != nil {
fail(c, http.StatusInternalServerError, err.Error())
return
}
ok(c, gin.H{"total": total, "items": users})
}
// Get 用户详情。
func (h *UserHandler) Get(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
fail(c, http.StatusBadRequest, "无效的用户 ID")
return
}
u, err := h.svc.GetByID(c.Request.Context(), uint(id))
if err != nil {
fail(c, http.StatusNotFound, err.Error())
return
}
ok(c, u)
}
// UserUpdateRequest 更新外部用户信息(仅更新提供的字段;邮箱仅管理员可改)。
type UserUpdateRequest struct {
Email *string `json:"email"`
Supervisor *string `json:"supervisor"`
Purpose *string `json:"purpose"`
}
// Update PATCH /users/:id。
func (h *UserHandler) Update(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
fail(c, http.StatusBadRequest, "无效的用户 ID")
return
}
var req UserUpdateRequest
if err := c.ShouldBindJSON(&req); err != nil {
fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error())
return
}
u, err := h.svc.Update(c.Request.Context(), uint(id), req.Email, req.Supervisor, req.Purpose)
if err != nil {
h.h.audit(c, "user.update", "user", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed)
if errors.Is(err, service.ErrUserNotFound) {
fail(c, http.StatusNotFound, err.Error())
return
}
fail(c, http.StatusBadRequest, err.Error())
return
}
h.h.audit(c, "user.update", "user", c.Param("id"), map[string]any{"email": req.Email, "supervisor": req.Supervisor, "purpose": req.Purpose}, model.ResultSuccess)
ok(c, u)
}
// Disable POST /users/:id/disable —— 禁用(清空 authorized_keysSSH 立即失效)。
func (h *UserHandler) Disable(c *gin.Context) {
h.setStatus(c, "user.disable", model.UserStatusDisabled, h.svc.Disable)
}
// Enable POST /users/:id/enable —— 启用(按 DB 密钥状态恢复)。
func (h *UserHandler) Enable(c *gin.Context) {
h.setStatus(c, "user.enable", model.UserStatusActive, h.svc.Enable)
}
// setStatus 复用禁用/启用的公共流程(解析 ID、调 service、审计)。
func (h *UserHandler) setStatus(c *gin.Context, action, wantStatus string, fn func(ctx context.Context, id uint) error) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
fail(c, http.StatusBadRequest, "无效的用户 ID")
return
}
if err := fn(c, uint(id)); err != nil {
h.h.audit(c, action, "user", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed)
switch {
case errors.Is(err, service.ErrUserNotFound):
fail(c, http.StatusNotFound, err.Error())
case errors.Is(err, service.ErrUserExpired):
fail(c, http.StatusConflict, err.Error())
case errors.Is(err, service.ErrSystemAccountMissing):
fail(c, http.StatusConflict, err.Error())
default:
fail(c, http.StatusInternalServerError, err.Error())
}
return
}
ok(c, gin.H{"id": u.ID, "username": u.Username, "status": u.Status})
h.h.audit(c, action, "user", c.Param("id"), map[string]any{"status": wantStatus}, model.ResultSuccess)
ok(c, gin.H{"status": wantStatus})
}
// List 用户列表(M1 实现分页筛选)
func (h *UserHandler) List(c *gin.Context) {
fail(c, http.StatusNotImplemented, "用户列表将在 M1 实现")
// ExtendRequest 延期请求
type ExtendRequest struct {
Days int `json:"days"` // 0 表示用配置默认(90 天)
}
// Extend POST /users/:id/extend —— 延长有效期;已过期用户在回收期内可恢复。
func (h *UserHandler) Extend(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
fail(c, http.StatusBadRequest, "无效的用户 ID")
return
}
var req ExtendRequest
if err := c.ShouldBindJSON(&req); err != nil {
fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error())
return
}
u, err := h.svc.GetByID(c.Request.Context(), uint(id))
if err != nil {
fail(c, http.StatusNotFound, err.Error())
return
}
if err := h.svc.Extend(c.Request.Context(), uint(id), req.Days); err != nil {
h.h.audit(c, "user.extend", "user", c.Param("id"), map[string]any{"days": req.Days, "err": err.Error()}, model.ResultFailed)
fail(c, http.StatusInternalServerError, err.Error())
return
}
h.h.audit(c, "user.extend", "user", c.Param("id"), map[string]any{"days": req.Days, "old_status": u.Status}, model.ResultSuccess)
ok(c, gin.H{"expire_at": time.Now().Add(h.extendTTL(req.Days)).UTC()})
}
// extendTTL 计算新的有效期(与 service 保持一致:days<=0 用默认)。
func (h *UserHandler) extendTTL(days int) time.Duration {
if days <= 0 {
return h.cfg.Policy.DefaultTTL
}
return time.Duration(days) * 24 * time.Hour
}
// Delete DELETE /users/:id —— 删除并回收(系统账号 + 家目录 + 密钥,保留审计)。
func (h *UserHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
fail(c, http.StatusBadRequest, "无效的用户 ID")
return
}
if err := h.svc.Delete(c.Request.Context(), uint(id)); err != nil {
h.h.audit(c, "user.delete", "user", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed)
switch {
case errors.Is(err, service.ErrUserNotFound):
fail(c, http.StatusNotFound, err.Error())
default:
fail(c, http.StatusInternalServerError, err.Error())
}
return
}
h.h.audit(c, "user.delete", "user", c.Param("id"), nil, model.ResultSuccess)
ok(c, gin.H{"status": "deleted"})
}
+7 -5
View File
@@ -14,10 +14,11 @@ var ErrCaptchaInvalid = errors.New("auth: 图形验证码错误")
// Captcha 图形验证码(防机器人,登录前置)。
type Captcha struct {
ID string
Text string // M1 生成图像渲染,此处仅存文本
Text string
}
// CaptchaStore 为图形验证码存储M1 实现图像渲染)。
// CaptchaStore 为图形验证码存储。单实例内存实现为默认;
// 多实例部署需改 DB/RedisPLAN §6 注明)。
type CaptchaStore interface {
// New 生成一个验证码并返回其 ID。
New() (*Captcha, error)
@@ -27,6 +28,7 @@ type CaptchaStore interface {
// MemoryCaptchaStore 单实例内存实现。
type MemoryCaptchaStore struct {
ttl time.Duration
mu sync.Mutex
entries map[string]*captchaEntry
}
@@ -37,8 +39,8 @@ type captchaEntry struct {
}
// NewMemoryCaptchaStore 创建内存图形验证码存储。
func NewMemoryCaptchaStore() *MemoryCaptchaStore {
return &MemoryCaptchaStore{entries: make(map[string]*captchaEntry)}
func NewMemoryCaptchaStore(ttl time.Duration) *MemoryCaptchaStore {
return &MemoryCaptchaStore{ttl: ttl, entries: make(map[string]*captchaEntry)}
}
func (s *MemoryCaptchaStore) New() (*Captcha, error) {
@@ -52,7 +54,7 @@ func (s *MemoryCaptchaStore) New() (*Captcha, error) {
}
s.mu.Lock()
defer s.mu.Unlock()
s.entries[id] = &captchaEntry{text: text, expiresAt: time.Now().Add(5 * time.Minute)}
s.entries[id] = &captchaEntry{text: text, expiresAt: time.Now().Add(s.ttl)}
return &Captcha{ID: id, Text: text}, nil
}
+131
View File
@@ -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
}
+115
View File
@@ -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
View File
@@ -1,15 +1,22 @@
// Package auth 提供认证相关能力:OTP(双通道)、会话、bcrypt、图形验证码。
//
// OTP 双通道对齐:邮件发送与 CLI 获取共用同一 OTPStore(同一验证码、同一
// 10 分钟有效期、同一 60s 冷却与失败限速),邮件失败不阻断 CLI 通道。
// 单实例用内存存储;多实例需改为 DB/Redis(PLAN §6 注明)。
// 有效期、同一冷却与失败限速),邮件失败不阻断 CLI 通道。邮件通道经
// Send 生成验证码,CLI 通道优先经 Current 复用同一验证码,无有效码时才
// 触发 Send(仍受同一冷却约束)。验证码落 DB(otp_codes 表),单实例部署
// 即可保证跨进程(HTTP 服务与 CLI 子命令)共享;多实例需改 DB 行锁/Redis
// PLAN §6)。
package auth
import (
"context"
"errors"
"sync"
"time"
"gorm.io/gorm"
"ws_usernode/internal/model"
"ws_usernode/internal/pkg"
)
@@ -21,20 +28,141 @@ var (
)
const (
otpCodeLen = 6
maxFailures = 5 // 单账号连续失败限速阈值
failureWin = 10 * time.Minute // 失败计数窗口
otpCodeLen = 6
)
// OTPStore 为 OTP 验证码存储。内存实现为单实例默认实现
// OTP 失败限速默认参数(单账号连续失败阈值与计数窗口)
const (
DefaultMaxFailures = 5
DefaultFailureWin = 10 * time.Minute
)
// OTPStore 为 OTP 验证码存储。DB 实现(DBOTPStore)为生产默认,
// MemoryOTPStore 供测试与单进程内嵌场景使用。
type OTPStore interface {
// Send 为 username 生成新验证码覆盖旧码)。冷却期内调用返回 ErrCooldown。
// 邮件与 CLI 双通道都走该方法,保证对齐。
Send(username string, ttl, cooldown time.Duration) (string, error)
// Send 为 username 生成新验证码覆盖旧码(邮件通道)。冷却期内返回 ErrCooldown。
Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error)
// Current 返回当前有效(未过期、未消费)验证码,供 CLI 通道复用同一验证码。
// 无有效验证码返回 ErrInvalidCode。
Current(ctx context.Context, username string) (string, error)
// Verify 校验验证码并一次性消费。失败累计计数(达到阈值返回 ErrTooManyFails)。
Verify(username, code string) (bool, error)
Verify(ctx context.Context, username, code string) (bool, error)
// Failures 返回 username 当前失败计数。
Failures(username string) (int, error)
Failures(ctx context.Context, username string) (int, error)
}
// DBOTPStore 基于 model.OTPCode 的存储实现,每用户一行(username 唯一)。
type DBOTPStore struct {
db *gorm.DB
maxFailures int
failureWin time.Duration
}
// NewDBOTPStore 创建 DB OTP 存储。
func NewDBOTPStore(db *gorm.DB, maxFailures int, failureWin time.Duration) *DBOTPStore {
return &DBOTPStore{db: db, maxFailures: maxFailures, failureWin: failureWin}
}
func (s *DBOTPStore) Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) {
now := time.Now()
var row model.OTPCode
err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error
// 注意:First 的 err 不能复用给后续语句,避免被覆盖导致走错分支
exists := err == nil
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return "", err
}
if exists && now.Before(row.CooldownUntil) {
return "", ErrCooldown
}
code, randErr := pkg.RandomDigits(otpCodeLen)
if randErr != nil {
return "", randErr
}
updates := map[string]any{
"code": code,
"expires_at": now.Add(ttl),
"cooldown_until": now.Add(cooldown),
"failures": 0,
"failed_at": nil,
"consumed_at": nil,
}
if exists {
err = s.db.WithContext(ctx).Model(&row).Updates(updates).Error
} else {
err = s.db.WithContext(ctx).Create(&model.OTPCode{
Username: username,
Code: code,
ExpiresAt: now.Add(ttl),
CooldownUntil: now.Add(cooldown),
}).Error
}
if err != nil {
return "", err
}
return code, nil
}
func (s *DBOTPStore) Current(ctx context.Context, username string) (string, error) {
var row model.OTPCode
if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", ErrInvalidCode
}
return "", err
}
if row.ConsumedAt != nil || time.Now().After(row.ExpiresAt) {
return "", ErrInvalidCode
}
return row.Code, nil
}
func (s *DBOTPStore) Verify(ctx context.Context, username, code string) (bool, error) {
now := time.Now()
var row model.OTPCode
if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, ErrInvalidCode
}
return false, err
}
if row.ConsumedAt != nil || now.After(row.ExpiresAt) {
return false, ErrInvalidCode
}
// 失败计数窗口:超出窗口则重置计数
if row.FailedAt != nil && now.Sub(*row.FailedAt) > s.failureWin {
row.Failures = 0
row.FailedAt = nil
}
if row.Failures >= s.maxFailures {
return false, ErrTooManyFails
}
if row.Code != code {
row.Failures++
f := now
row.FailedAt = &f
_ = s.db.WithContext(ctx).Model(&row).Updates(map[string]any{"failures": row.Failures, "failed_at": row.FailedAt}).Error
if row.Failures >= s.maxFailures {
return false, ErrTooManyFails
}
return false, ErrInvalidCode
}
consumed := now
if err := s.db.WithContext(ctx).Model(&row).Update("consumed_at", &consumed).Error; err != nil {
return false, err
}
return true, nil
}
func (s *DBOTPStore) Failures(ctx context.Context, username string) (int, error) {
var row model.OTPCode
if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return 0, nil
}
return 0, err
}
return row.Failures, nil
}
type otpEntry struct {
@@ -44,7 +172,7 @@ type otpEntry struct {
failures int
}
// MemoryOTPStore 为单实例内存实现。
// MemoryOTPStore 为单进程内存实现(测试/内嵌场景)
type MemoryOTPStore struct {
mu sync.Mutex
entries map[string]*otpEntry
@@ -55,9 +183,10 @@ func NewMemoryOTPStore() *MemoryOTPStore {
return &MemoryOTPStore{entries: make(map[string]*otpEntry)}
}
func (s *MemoryOTPStore) Send(username string, ttl, cooldown time.Duration) (string, error) {
func (s *MemoryOTPStore) Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
now := time.Now()
if e, ok := s.entries[username]; ok && now.Before(e.cooldownAt) {
@@ -76,9 +205,21 @@ func (s *MemoryOTPStore) Send(username string, ttl, cooldown time.Duration) (str
return code, nil
}
func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
func (s *MemoryOTPStore) Current(ctx context.Context, username string) (string, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
e, ok := s.entries[username]
if !ok || time.Now().After(e.expiresAt) {
return "", ErrInvalidCode
}
return e.code, nil
}
func (s *MemoryOTPStore) Verify(ctx context.Context, username, code string) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
e, ok := s.entries[username]
if !ok {
@@ -89,12 +230,12 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
delete(s.entries, username)
return false, ErrInvalidCode
}
if e.failures >= maxFailures {
if e.failures >= DefaultMaxFailures {
return false, ErrTooManyFails
}
if e.code != code {
e.failures++
if e.failures >= maxFailures {
if e.failures >= DefaultMaxFailures {
return false, ErrTooManyFails
}
return false, ErrInvalidCode
@@ -103,9 +244,10 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
return true, nil
}
func (s *MemoryOTPStore) Failures(username string) (int, error) {
func (s *MemoryOTPStore) Failures(ctx context.Context, username string) (int, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
if e, ok := s.entries[username]; ok {
return e.failures, nil
}
+119
View File
@@ -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)
}
}
+68
View File
@@ -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)
}
+76
View File
@@ -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
View File
@@ -28,13 +28,15 @@ type Config struct {
Database DatabaseConfig `toml:"database"`
Log LogConfig `toml:"log"`
Policy PolicyConfig `toml:"policy"`
Auth AuthConfig `toml:"auth"`
SMTP SMTPConfig `toml:"smtp"`
System SystemConfig `toml:"system"`
}
type AppConfig struct {
Name string `toml:"name"`
Env string `toml:"env"` // development / production
Name string `toml:"name"`
Env string `toml:"env"` // development / production
BaseURL string `toml:"base_url"` // 对外访问地址(邮件中的链接使用)
}
type ServerConfig struct {
@@ -62,6 +64,13 @@ type PolicyConfig struct {
OTPCooldown time.Duration `toml:"otp_cooldown"` // OTP 发送冷却
}
// AuthConfig 认证与防暴力参数。
type AuthConfig struct {
MaxLoginFailures int `toml:"max_login_failures"` // 管理员登录连续失败阈值,达到后锁定
LockDuration time.Duration `toml:"lock_duration"` // 失败达到阈值后的锁定时长
CaptchaTTL time.Duration `toml:"captcha_ttl"` // 图形验证码有效期
}
type SMTPConfig struct {
Host string `toml:"host"`
Port int `toml:"port"`
@@ -72,21 +81,22 @@ type SMTPConfig struct {
// SystemConfig 为系统账号操作层的本地实现配置(sudoers 白名单模式)。
type SystemConfig struct {
Sudo bool `toml:"sudo"` // 是否通过 sudo -n 执行系统命令;开发环境 false 时 dry-run
UserPrefix string `toml:"user_prefix"` // 外部用户系统账号前缀,默认 ext_
Group string `toml:"group"` // 外部用户所属组,默认 external
Shell string `toml:"shell"` // 默认 shell
HomeBase string `toml:"home_base"` // 家目录基路径
AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖)
Sudo bool `toml:"sudo"` // 是否通过 sudo -n 执行系统命令(生产)
DryRun bool `toml:"dry_run"` // true = 只打印计划命令不执行(开发演练);false 且 sudo=false 时直接执行(容器/测试用户验证)
UserPrefix string `toml:"user_prefix"` // 外部用户系统账号前缀,默认 ext_
Group string `toml:"group"` // 外部用户所属组,默认 external
Shell string `toml:"shell"` // 默认 shell
HomeBase string `toml:"home_base"` // 家目录基路径
AuthorizedKeysDir string `toml:"authorized_keys_dir"` // authorized_keys 所在目录(测试可覆盖)
}
// Default 返回带开发环境默认值的配置,作为 config.example.toml 与未配置项的兜底。
func Default() *Config {
return &Config{
App: AppConfig{Name: "ws_usernode", Env: "development"},
App: AppConfig{Name: "ws_usernode", Env: "development", BaseURL: "http://127.0.0.1:8080"},
Server: ServerConfig{
Listen: "127.0.0.1:8080",
SessionTTL: 24 * time.Hour,
Listen: "127.0.0.1:8080",
SessionTTL: 24 * time.Hour,
TrustedProxies: []string{"127.0.0.1", "::1"},
},
Database: DatabaseConfig{Driver: "sqlite", DSN: "data/usernode.db"},
@@ -98,9 +108,17 @@ func Default() *Config {
OTPTTL: 10 * time.Minute,
OTPCooldown: 60 * time.Second,
},
Auth: AuthConfig{
MaxLoginFailures: 5,
LockDuration: 15 * time.Minute,
CaptchaTTL: 5 * time.Minute,
},
SMTP: SMTPConfig{Port: 587},
System: SystemConfig{
// 开发默认 dry-run:未配置 config 直接跑 serve 时只打印计划,避免误操作系统账号。
// 生产必须显式 dry_run=false 且 sudo=true(见 deploy/sudoers.example)。
Sudo: false,
DryRun: true,
UserPrefix: "ext_",
Group: "external",
Shell: "/bin/sh",
+16
View File
@@ -31,6 +31,12 @@ func TestLoadDefault(t *testing.T) {
if cfg.System.UserPrefix != "ext_" {
t.Errorf("default user_prefix = %q", cfg.System.UserPrefix)
}
if cfg.Auth.MaxLoginFailures != 5 || cfg.Auth.LockDuration != 15*time.Minute {
t.Errorf("default auth = %+v", cfg.Auth)
}
if !cfg.System.DryRun {
t.Error("default dry_run should be true (安全默认)")
}
}
func TestLoadFileOverrides(t *testing.T) {
@@ -62,7 +68,11 @@ func TestEnvOverrides(t *testing.T) {
t.Setenv("USERNODE_DATABASE_DSN", "u:p@tcp(h:3306)/db")
t.Setenv("USERNODE_POLICY_OTPTTL", "5m")
t.Setenv("USERNODE_SYSTEM_SUDO", "true")
t.Setenv("USERNODE_SYSTEM_DRY_RUN", "false")
t.Setenv("USERNODE_SERVER_TRUSTED_PROXIES", "10.0.0.1, 10.0.0.2")
t.Setenv("USERNODE_AUTH_MAX_LOGIN_FAILURES", "3")
t.Setenv("USERNODE_AUTH_LOCK_DURATION", "5m")
t.Setenv("USERNODE_APP_BASE_URL", "https://un.example.com")
cfg, err := LoadDefault()
if err != nil {
@@ -77,6 +87,12 @@ func TestEnvOverrides(t *testing.T) {
if !cfg.System.Sudo {
t.Error("system.sudo should be true")
}
if cfg.Auth.MaxLoginFailures != 3 || cfg.Auth.LockDuration != 5*time.Minute {
t.Errorf("auth = %+v", cfg.Auth)
}
if cfg.App.BaseURL != "https://un.example.com" {
t.Errorf("base_url = %q", cfg.App.BaseURL)
}
if len(cfg.Server.TrustedProxies) != 2 || cfg.Server.TrustedProxies[0] != "10.0.0.1" {
t.Errorf("trusted_proxies = %v", cfg.Server.TrustedProxies)
}
+118
View File
@@ -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()
}
+32
View File
@@ -141,11 +141,43 @@ type MailLog struct {
UpdatedAt time.Time `json:"updated_at"`
}
// OTPCode 外部用户 OTP 验证码(DB 存储,邮件与 CLI 双通道共享)。
//
// 每用户一行(username 唯一):邮件通道经 Send 生成,CLI 通道经 Current
// 读取同一验证码,保证"同一验证码、同一有效期、同一冷却与失败限速"(PLAN §2.2)。
// 验证码为 6 位数字,生命周期短(10 分钟)且一次性消费,存储明文以便 CLI
// 复用返回;并发写由单实例部署的串行事务保证(多实例需改 DB 行锁/Redis,PLAN §6)。
type OTPCode struct {
ID uint `gorm:"primaryKey" json:"id"`
Username string `gorm:"size:64;uniqueIndex;not null" json:"username"`
Code string `gorm:"size:16;not null" json:"-"`
ExpiresAt time.Time `gorm:"index;not null" json:"expires_at"`
CooldownUntil time.Time `json:"cooldown_until"` // Send 冷却截止
Failures int `json:"failures"`
FailedAt *time.Time `json:"failed_at"` // 失败计数窗口起点(窗口内达阈值限速)
ConsumedAt *time.Time `json:"consumed_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// PasswordResetToken 管理员密码重置令牌(邮件重置)。DB 只存哈希,
// 明文令牌仅经邮件/日志发送给管理员,一次有效。
type PasswordResetToken struct {
ID uint `gorm:"primaryKey" json:"id"`
AdminID uint `gorm:"index;not null" json:"admin_id"`
TokenHash string `gorm:"size:64;not null" json:"-"`
ExpiresAt time.Time `gorm:"index;not null" json:"expires_at"`
UsedAt *time.Time `json:"used_at"`
IP string `gorm:"size:64" json:"ip"`
CreatedAt time.Time `json:"created_at"`
}
// AllModels 供 AutoMigrate 使用的全部模型。
func AllModels() []any {
return []any{
&AdminUser{}, &User{}, &SSHKey{}, &Approval{},
&AuditLog{}, &Session{}, &Setting{}, &MailLog{},
&OTPCode{}, &PasswordResetToken{},
}
}
+46
View File
@@ -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
View File
@@ -9,13 +9,14 @@ import (
"github.com/gin-gonic/gin"
"ws_usernode/internal/api"
"ws_usernode/internal/auth"
"ws_usernode/internal/config"
"ws_usernode/internal/webui"
)
// New 构建根 routerAPI v1 + 前端静态资源(go:embed)。
// production 模式启用 gin.ReleaseMode;否则启用调试模式与开发日志。
func New(cfg *config.Config, h *api.Handler, log *slog.Logger) *gin.Engine {
func New(cfg *config.Config, h *api.Handler, sessions auth.SessionStore, log *slog.Logger) *gin.Engine {
if cfg.App.Env == "production" {
gin.SetMode(gin.ReleaseMode)
}
@@ -28,21 +29,34 @@ func New(cfg *config.Config, h *api.Handler, log *slog.Logger) *gin.Engine {
// RESTful API v1
v1 := r.Group("/api/v1")
{
auth := v1.Group("/auth")
authGrp := v1.Group("/auth")
{
// M1captcha / otp/send / otp/login / admin/login / logout / me
auth.GET("/captcha", notImplemented("图形验证码(M1"))
auth.POST("/otp/send", notImplemented("OTP 发送(M1"))
auth.POST("/otp/login", notImplemented("OTP 登录(M1"))
auth.POST("/admin/login", notImplemented("管理员登录(M1"))
auth.POST("/logout", notImplemented("登出(M1"))
auth.GET("/me", notImplemented("当前会话(M1"))
authGrp.GET("/captcha", h.Auth.Captcha)
authGrp.POST("/otp/send", h.Auth.OTPSend)
authGrp.POST("/otp/login", h.Auth.OTPLogin)
authGrp.POST("/admin/login", h.Auth.AdminLogin)
authGrp.POST("/admin/forgot", h.Auth.AdminForgot)
authGrp.POST("/admin/reset", h.Auth.AdminReset)
// 需要会话(管理员或外部用户)
authed := authGrp.Group("", sessionMiddleware(sessions))
authed.POST("/logout", h.Auth.Logout)
authed.GET("/me", h.Auth.Me)
}
v1.GET("/users", h.User.List)
v1.POST("/users", h.User.Create)
v1.POST("/users/:id/disable", notImplemented("禁用用户(M1"))
v1.POST("/users/:id/enable", notImplemented("启用用户(M1"))
v1.POST("/users/:id/extend", notImplemented("延期(M1"))
// 用户管理(admin
users := v1.Group("/users", sessionMiddleware(sessions), requireUserType(auth.SessionUserAdmin))
{
users.GET("", h.User.List)
users.POST("", h.User.Create)
users.GET("/:id", h.User.Get)
users.PATCH("/:id", h.User.Update)
users.POST("/:id/disable", h.User.Disable)
users.POST("/:id/enable", h.User.Enable)
users.POST("/:id/extend", h.User.Extend)
users.DELETE("/:id", h.User.Delete)
}
// 后续里程碑
v1.POST("/approvals", notImplemented("提交申请(M3"))
v1.GET("/approvals", notImplemented("申请列表(M3"))
v1.POST("/approvals/:id/review", notImplemented("审批(M3"))
+41 -3
View File
@@ -15,9 +15,10 @@ import (
// 管理员服务错误。
var (
ErrAdminExists = errors.New("service: 管理员已存在")
ErrAdminNotFound = errors.New("service: 管理员不存在")
ErrWeakPassword = errors.New("service: 密码过弱(至少 8 位,需含字母与数字)")
ErrAdminExists = errors.New("service: 管理员已存在")
ErrAdminNotFound = errors.New("service: 管理员不存在")
ErrWeakPassword = errors.New("service: 密码过弱(至少 8 位,需含字母与数字)")
ErrBadCredentials = errors.New("service: 用户名或密码错误")
)
// AdminService 管理端账号服务(CLI admin create / reset-password 与登录共用)。
@@ -94,6 +95,43 @@ func (s *AdminService) GetByUsername(ctx context.Context, username string) (*mod
return &adm, nil
}
// GetByID 按 ID 查询管理员。
func (s *AdminService) GetByID(ctx context.Context, id uint) (*model.AdminUser, error) {
var adm model.AdminUser
if err := s.db.WithContext(ctx).First(&adm, "id = ?", id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrAdminNotFound
}
return nil, err
}
return &adm, nil
}
// GetByEmail 按邮箱查询管理员(用于密码重置邮件)。
func (s *AdminService) GetByEmail(ctx context.Context, email string) (*model.AdminUser, error) {
var adm model.AdminUser
if err := s.db.WithContext(ctx).First(&adm, "email = ?", email).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrAdminNotFound
}
return nil, err
}
return &adm, nil
}
// Login 校验管理员用户名与口令(bcrypt)。失败统一返回 ErrBadCredentials
// 不暴露账号是否存在。登录限速由 AuthService 的 RateLimiter 处理。
func (s *AdminService) Login(ctx context.Context, username, password string) (*model.AdminUser, error) {
adm, err := s.GetByUsername(ctx, strings.TrimSpace(username))
if err != nil {
return nil, ErrBadCredentials
}
if !auth.VerifyPassword(adm.PasswordHash, password) {
return nil, ErrBadCredentials
}
return adm, nil
}
// validatePassword 校验管理员密码强度(骨架阶段基础规则,M1 可加策略)。
func validatePassword(p string) error {
if len(p) < 8 {
+235
View File
@@ -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: 会话无效")
+230
View File
@@ -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
View File
@@ -4,9 +4,11 @@ import (
"context"
"errors"
"strings"
"time"
"gorm.io/gorm"
"ws_usernode/internal/config"
"ws_usernode/internal/model"
"ws_usernode/internal/pkg"
"ws_usernode/internal/system"
@@ -14,29 +16,30 @@ import (
// 用户服务错误。
var (
ErrUserNotFound = errors.New("service: 用户不存在")
ErrUserExists = errors.New("service: 用户名已存在")
ErrUserNotFound = errors.New("service: 用户不存在")
ErrUserExists = errors.New("service: 用户名已存在")
ErrUserExpired = errors.New("service: 用户已过期,请先延期")
ErrUserDisabled = errors.New("service: 用户已禁用,无法操作")
ErrSystemAccountMissing = errors.New("service: 系统账号不存在,无法操作")
)
// UserService 外部用户生命周期服务。
// M0 提供查询与创建骨架;建号(useradd)、禁用/延期等系统操作 M1 接入 system.Manager。
// UserService 外部用户生命周期服务DB 记录 + system.Manager 系统账号操作
type UserService struct {
db *gorm.DB
sys system.Manager
cfg *config.Config
}
// NewUserService 创建用户服务。
func NewUserService(db *gorm.DB, sys system.Manager) *UserService {
return &UserService{db: db, sys: sys}
func NewUserService(db *gorm.DB, sys system.Manager, cfg *config.Config) *UserService {
return &UserService{db: db, sys: sys, cfg: cfg}
}
// GetByUsername 按用户名查询外部用户(含或不含 ext_ 前缀均可)。
func (s *UserService) GetByUsername(ctx context.Context, username string) (*model.User, error) {
if !strings.HasPrefix(username, "ext_") {
username = "ext_" + username
}
full := normalizeName(username, s.cfg.System.UserPrefix)
var u model.User
if err := s.db.WithContext(ctx).First(&u, "username = ?", username).Error; err != nil {
if err := s.db.WithContext(ctx).First(&u, "username = ?", full).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrUserNotFound
}
@@ -45,16 +48,68 @@ func (s *UserService) GetByUsername(ctx context.Context, username string) (*mode
return &u, nil
}
// Create 创建外部用户记录并调用系统层建号。M0 阶段系统层为 dry-run
// username 为不含前缀的申请名,内部加 ext_ 前缀。
func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttlSeconds int64) (*model.User, error) {
// GetByID 按 ID 查询外部用户
func (s *UserService) GetByID(ctx context.Context, id uint) (*model.User, error) {
var u model.User
if err := s.db.WithContext(ctx).First(&u, "id = ?", id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrUserNotFound
}
return nil, err
}
return &u, nil
}
// UserFilter 用户列表筛选条件。
type UserFilter struct {
Status string // active / disabled / expired,空为全部
Supervisor string // 挂靠老师模糊匹配
Page int
PageSize int
}
// List 分页查询用户(admin)。
func (s *UserService) List(ctx context.Context, f UserFilter) ([]model.User, int64, error) {
q := s.db.WithContext(ctx).Model(&model.User{})
if f.Status != "" {
q = q.Where("status = ?", f.Status)
}
if f.Supervisor != "" {
q = q.Where("supervisor LIKE ?", "%"+f.Supervisor+"%")
}
var total int64
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
page, size := f.Page, f.PageSize
if page < 1 {
page = 1
}
if size < 1 {
size = 20
}
if size > 100 {
size = 100
}
var users []model.User
if err := q.Order("id DESC").Offset((page - 1) * size).Limit(size).Find(&users).Error; err != nil {
return nil, 0, err
}
return users, total, nil
}
// Create 创建外部用户:DB 记录 + 系统账号(useradd + passwd -l)。
// username 不含前缀;ttl 为有效期时长,0 表示用配置默认(90 天)。
// 系统建号失败时回滚 DB 记录,保证两侧一致。
func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttl time.Duration, createdBy uint) (*model.User, error) {
if err := pkg.ValidateUserName(username); err != nil {
return nil, err
}
if err := pkg.ValidateEmail(email); err != nil {
return nil, err
}
full := "ext_" + username
username = strings.TrimSpace(username)
full := s.cfg.System.UserPrefix + username
var count int64
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", full).Count(&count).Error; err != nil {
return nil, err
@@ -62,21 +117,178 @@ func (s *UserService) Create(ctx context.Context, username, email, supervisor, p
if count > 0 {
return nil, ErrUserExists
}
if ttl <= 0 {
ttl = s.cfg.Policy.DefaultTTL
}
expireAt := time.Now().Add(ttl)
u := &model.User{
Username: full,
Email: email,
Supervisor: supervisor,
Purpose: purpose,
Status: model.UserStatusActive,
Shell: "/bin/sh",
ExpireAt: &expireAt,
Shell: s.cfg.System.Shell,
CreatedBy: createdBy,
}
if err := s.db.WithContext(ctx).Create(u).Error; err != nil {
return nil, err
}
// 系统账号创建(dry-run / 真实),失败时回滚 DB 记录
if err := s.sys.CreateUser(ctx, system.Account{Username: full}); err != nil {
// 系统账号创建(dry-run / 直接 / sudo),失败时回滚 DB 记录
if err := s.sys.CreateUser(ctx, system.Account{Username: full, Shell: s.cfg.System.Shell}); err != nil {
_ = s.db.WithContext(ctx).Delete(u).Error
return nil, err
}
return u, nil
}
// Update 更新外部用户信息(仅更新非 nil 字段;邮箱由管理员修改,用户不可自助改)。
func (s *UserService) Update(ctx context.Context, id uint, email, supervisor, purpose *string) (*model.User, error) {
if _, err := s.GetByID(ctx, id); err != nil {
return nil, err
}
updates := make(map[string]any)
if email != nil {
if err := pkg.ValidateEmail(*email); err != nil {
return nil, err
}
updates["email"] = *email
}
if supervisor != nil {
updates["supervisor"] = *supervisor
}
if purpose != nil {
updates["purpose"] = *purpose
}
if len(updates) > 0 {
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err
}
}
return s.GetByID(ctx, id)
}
// systemAccountOK 检查系统账号存在;dry-run 模式跳过检查(演练流程)。
func (s *UserService) systemAccountOK(ctx context.Context, username string) bool {
if s.cfg.System.DryRun {
return true
}
ok, err := s.sys.Exists(ctx, username)
return err == nil && ok
}
// Disable 禁用用户:DB 置 disabled + 清空 authorized_keysSSH 立即失效)。
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_atdays<=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
}
+242
View File
@@ -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
View File
@@ -4,15 +4,19 @@
// 执行 useradd/usermod/userdel/passwd 等固定命令并做参数强校验。未来多节点
// agent 模式只需新增远程实现替换本地实现(PLAN §5.1)。
//
// 权限模型:节点以专有用户(如 usernode)运行,经 sudo -n 提权执行白名单
// 命令;开发环境 config system.sudo=false 进入 dry-run只打印计划不执行),
// 避免在开发机上直接操作系统账号。系统命令的真实系统效果在 M1 用测试用户/
// 容器验证,不在生产直接跑 useradd。
// 执行模式(config system.*):
// - dry_run=true只打印计划命令不执行(开发演练);
// - dry_run=false + sudo=false:直接执行(容器/测试用户验证真实建号);
// - dry_run=false + sudo=true:经 sudo -n 提权执行白名单命令(生产,
// 需配置 deploy/sudoers)。真实系统账号操作只在测试用户/容器中验证,
// 不在生产直接跑 useradd。
package system
import (
"bufio"
"context"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
@@ -45,8 +49,10 @@ type Manager interface {
RemoveUser(ctx context.Context, username string) error
// SetLock 锁定/解锁系统账号口令(passwd -l / -u)。
SetLock(ctx context.Context, username string, locked bool) error
// Exists 检查系统账号是否存在(读 /etc/passwd,无需提权)。
Exists(ctx context.Context, username string) (bool, error)
// SyncAuthorizedKeys 以 DB 状态全量重写 authorized_keys(原子写 + 并发锁),
// 吊销密钥即从文件移除、立即失效。M0 提供 dry-run 实现,M2 完成生产路径。
// 吊销密钥即从文件移除、立即失效。
SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error
}
@@ -69,20 +75,32 @@ func New(cfg config.SystemConfig) Manager {
return &localManager{cfg: cfg}
}
// run 执行白名单命令sudo -n <cmd> <args...>参数在调用处强校验。
// dry-run 模式下只返回将执行的命令文本,不真正执行
// run 执行白名单命令参数在调用处强校验。
// 模式:dry-run 只返回命令文本;directsudo=false)直接执行
// sudo 经 sudo -n 提权执行(对应 deploy/sudoers 白名单)。
func (m *localManager) run(ctx context.Context, cmd string, args ...string) (string, error) {
if !allowedCommands[cmd] {
return "", errors.New("system: command not allowed: " + cmd)
}
argv := append([]string{"-n", cmd}, args...)
cmdline := strings.Join(append([]string{"sudo", "-n", cmd}, args...), " ")
if !m.cfg.Sudo {
return cmdline, nil // dry-run
cmdline := strings.Join(append([]string{cmd}, args...), " ")
if m.cfg.DryRun {
return cmdline, nil
}
var (
out []byte
err error
)
if m.cfg.Sudo {
cmdline = "sudo -n " + cmdline
out, err = exec.CommandContext(ctx, "sudo", append([]string{"-n", cmd}, args...)...).CombinedOutput()
} else {
out, err = exec.CommandContext(ctx, cmd, args...).CombinedOutput()
}
out, err := exec.CommandContext(ctx, "sudo", argv...).CombinedOutput()
if err != nil {
return string(out), err
if len(out) > 0 {
return cmdline, fmt.Errorf("%s: %s: %w", cmdline, strings.TrimSpace(string(out)), err)
}
return cmdline, fmt.Errorf("%s: %w", cmdline, err)
}
return cmdline, nil
}
@@ -109,11 +127,11 @@ func (m *localManager) CreateUser(ctx context.Context, acc Account) error {
home = filepath.Join(m.cfg.HomeBase, username)
}
if _, err := m.run(ctx, "useradd", "-m", "-d", home, "-s", shell, "-g", m.cfg.Group, username); err != nil {
return errors.New("system: useradd: " + err.Error())
return err
}
// 锁定口令,仅密钥登录
if _, err := m.run(ctx, "passwd", "-l", username); err != nil {
return errors.New("system: passwd -l: " + err.Error())
return err
}
return nil
}
@@ -124,7 +142,7 @@ func (m *localManager) RemoveUser(ctx context.Context, username string) error {
return err
}
if _, err := m.run(ctx, "userdel", "-r", username); err != nil {
return errors.New("system: userdel: " + err.Error())
return err
}
return nil
}
@@ -139,20 +157,51 @@ func (m *localManager) SetLock(ctx context.Context, username string, locked bool
flag = "-l"
}
if _, err := m.run(ctx, "passwd", flag, username); err != nil {
return errors.New("system: passwd: " + err.Error())
return err
}
return nil
}
// passwdPath 为系统账号数据库路径(测试可覆盖为临时文件)。
var passwdPath = "/etc/passwd"
// Exists 检查系统账号是否存在。直接读 /etc/passwd(世界可读,无需提权)。
func (m *localManager) Exists(ctx context.Context, username string) (bool, error) {
username = m.sysName(username)
if err := pkg.ValidateSystemAccount(username); err != nil {
return false, err
}
_ = ctx
f, err := os.Open(passwdPath)
if err != nil {
return false, err
}
defer f.Close()
sc := bufio.NewScanner(f)
for sc.Scan() {
line := sc.Text()
if line == "" {
continue
}
fields := strings.SplitN(line, ":", 2)
if len(fields) > 0 && fields[0] == username {
return true, nil
}
}
if err := sc.Err(); err != nil {
return false, err
}
return false, nil
}
// SyncAuthorizedKeys 全量重写 authorized_keys
//
// 1. 以 <lockBase>/<username>.lock 文件锁串行化并发写(flock);
// 2. 写临时文件(0600),再 rename 原子替换;
// 3. 无有效密钥时写入空文件(SSH 行为一致,避免文件缺失导致的歧义)。
//
// 生产模式(sudo=true)下节点进程具备对家目录 .ssh 的写权限——部署时经
// sudoers 白名单授予固定命令/受控脚本实现(M2 细化);当前实现直接做文件
// 操作并假设权限已配置,dry-run 模式打印计划命令。
// dry-run 模式只打印计划;真实写路径(direct/sudo)需要进程具备对家目录
// .ssh 的写权限——生产部署经 sudoers 白名单授予受控脚本实现(M2 细化)
func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error {
username = m.sysName(username)
if err := pkg.ValidateSystemAccount(username); err != nil {
@@ -169,13 +218,13 @@ func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string,
}
}
if !m.cfg.Sudo {
if m.cfg.DryRun {
var plan strings.Builder
plan.WriteString("mkdir -p " + sshDir + " (0700)\n")
plan.WriteString("flock " + lockPath + "\n")
plan.WriteString("write " + filepath.Join(sshDir, "authorized_keys") + " (0600)\n")
plan.WriteString(content.String())
// 骨架阶段:仅日志输出计划,不落盘
// 演练模式:仅日志输出计划,不落盘
return nil
}
+87
View File
@@ -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)
}
}