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:
@@ -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 {
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"ws_usernode/internal/auth"
|
||||
"ws_usernode/internal/config"
|
||||
"ws_usernode/internal/mail"
|
||||
"ws_usernode/internal/model"
|
||||
)
|
||||
|
||||
// 认证服务错误。
|
||||
var (
|
||||
ErrRateLimited = errors.New("service: 尝试次数过多,请稍后再试")
|
||||
ErrCaptchaFailed = errors.New("service: 图形验证码错误")
|
||||
ErrUserUnavailable = errors.New("service: 用户不存在或不可用")
|
||||
)
|
||||
|
||||
// AuthService 认证服务:管理员/外部用户登录、会话、图形验证码、OTP 双通道、
|
||||
// 密码重置。持有各存储与邮件发送器,作为 api 层与底层存储之间的桥。
|
||||
type AuthService struct {
|
||||
db *gorm.DB
|
||||
cfg *config.Config
|
||||
otps auth.OTPStore
|
||||
captchas auth.CaptchaStore
|
||||
sessions auth.SessionStore
|
||||
resets auth.ResetTokenStore
|
||||
limiter *auth.RateLimiter
|
||||
mailer mail.Mailer
|
||||
users *UserService
|
||||
admins *AdminService
|
||||
audit *AuditService
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
// NewAuthService 组装认证服务。
|
||||
func NewAuthService(db *gorm.DB, cfg *config.Config, otps auth.OTPStore, captchas auth.CaptchaStore,
|
||||
sessions auth.SessionStore, resets auth.ResetTokenStore, limiter *auth.RateLimiter,
|
||||
mailer mail.Mailer, users *UserService, admins *AdminService, audit *AuditService, log *slog.Logger) *AuthService {
|
||||
return &AuthService{
|
||||
db: db, cfg: cfg, otps: otps, captchas: captchas, sessions: sessions,
|
||||
resets: resets, limiter: limiter, mailer: mailer, users: users,
|
||||
admins: admins, audit: audit, log: log,
|
||||
}
|
||||
}
|
||||
|
||||
// NewCaptcha 生成图形验证码并渲染 PNG 图像。
|
||||
func (s *AuthService) NewCaptcha() (id string, png []byte, err error) {
|
||||
cap, err := s.captchas.New()
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
png, err = auth.RenderCaptchaPNG(cap.Text)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return cap.ID, png, nil
|
||||
}
|
||||
|
||||
// createSession 建立 cookie 会话(DB 存储)。
|
||||
func (s *AuthService) createSession(ctx context.Context, userType string, refID uint, ip, userAgent string) (string, error) {
|
||||
return s.sessions.Create(ctx, userType, refID, s.cfg.Server.SessionTTL, ip, userAgent)
|
||||
}
|
||||
|
||||
// AdminLogin 管理员用户名+口令登录,返回会话 ID。
|
||||
// 连续失败达到阈值(config auth.max_login_failures)后锁定 lock_duration。
|
||||
func (s *AuthService) AdminLogin(ctx context.Context, username, password, ip, userAgent string) (string, error) {
|
||||
key := "admin-login:" + strings.TrimSpace(username)
|
||||
if !s.limiter.Allow(key) {
|
||||
_ = s.audit.Record(ctx, 0, username, "admin.login", "admin", "", map[string]any{"locked": true}, ip, model.ResultFailed)
|
||||
return "", ErrRateLimited
|
||||
}
|
||||
adm, err := s.admins.Login(ctx, username, password)
|
||||
if err != nil {
|
||||
s.limiter.RecordFailure(key)
|
||||
_ = s.audit.Record(ctx, 0, username, "admin.login", "admin", "", nil, ip, model.ResultFailed)
|
||||
return "", ErrBadCredentials
|
||||
}
|
||||
s.limiter.Reset(key)
|
||||
sid, err := s.createSession(ctx, auth.SessionUserAdmin, adm.ID, ip, userAgent)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
_ = s.audit.Record(ctx, adm.ID, adm.Username, "admin.login", "admin", "", nil, ip, model.ResultSuccess)
|
||||
return sid, nil
|
||||
}
|
||||
|
||||
// Logout 注销会话(管理员或外部用户通用)。
|
||||
func (s *AuthService) Logout(ctx context.Context, sessionID string) error {
|
||||
return s.sessions.Delete(ctx, sessionID)
|
||||
}
|
||||
|
||||
// AdminForgot 发送密码重置邮件(无 SMTP 时退化为日志输出)。
|
||||
// 用户不存在时也返回成功,避免账号枚举。
|
||||
func (s *AuthService) AdminForgot(ctx context.Context, username, ip string) error {
|
||||
username = strings.TrimSpace(username)
|
||||
adm, err := s.admins.GetByUsername(ctx, username)
|
||||
if err != nil {
|
||||
s.log.Warn("admin forgot: user not found (not revealed)", "username", username)
|
||||
return nil
|
||||
}
|
||||
token, err := s.resets.Create(ctx, adm.ID, 30*time.Minute, ip)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link := strings.TrimRight(s.cfg.App.BaseURL, "/") + "/reset?token=" + token
|
||||
body := "您正在重置管理员密码。请在 30 分钟内打开以下链接完成重置:\n\n" + link +
|
||||
"\n\n如非本人操作请忽略本邮件。也可由管理员通过 CLI `usernode admin reset-password` 重置。"
|
||||
if err := s.mailer.Send(ctx, adm.Email, "重置密码", body); err != nil {
|
||||
s.log.Warn("admin forgot: mail send failed", "err", err)
|
||||
}
|
||||
_ = s.audit.Record(ctx, adm.ID, adm.Username, "admin.forgot", "admin", "", nil, ip, model.ResultSuccess)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AdminReset 通过令牌重置管理员密码,并使该管理员既有会话全部失效。
|
||||
func (s *AuthService) AdminReset(ctx context.Context, token, newPassword string, ip string) error {
|
||||
adminID, err := s.resets.Consume(ctx, token)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
adm, err := s.admins.GetByID(ctx, adminID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validatePassword(newPassword); err != nil {
|
||||
return err
|
||||
}
|
||||
hash, err := auth.HashPassword(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Model(&model.AdminUser{}).Where("id = ?", adm.ID).Update("password_hash", hash).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// 重置后吊销该管理员全部会话
|
||||
_ = s.db.WithContext(ctx).Where("user_type = ? AND ref_id = ?", auth.SessionUserAdmin, adm.ID).Delete(&model.Session{}).Error
|
||||
_ = s.audit.Record(ctx, adm.ID, adm.Username, "admin.reset", "admin", "", nil, ip, model.ResultSuccess)
|
||||
return nil
|
||||
}
|
||||
|
||||
// UserOTPSend 外部用户申请 OTP:图形验证码前置,生成验证码并邮件发送。
|
||||
// 邮件失败不阻断(CLI 通道兜底)。冷却/失败限速与 CLI 通道共享同一存储。
|
||||
func (s *AuthService) UserOTPSend(ctx context.Context, username, captchaID, captchaAnswer, ip string) error {
|
||||
if !s.captchas.Verify(captchaID, captchaAnswer) {
|
||||
return ErrCaptchaFailed
|
||||
}
|
||||
u, err := s.users.GetByUsername(ctx, username)
|
||||
if err != nil {
|
||||
return ErrUserUnavailable
|
||||
}
|
||||
if u.Status != model.UserStatusActive {
|
||||
return ErrUserUnavailable
|
||||
}
|
||||
code, err := s.otps.Send(ctx, u.Username, s.cfg.Policy.OTPTTL, s.cfg.Policy.OTPCooldown)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
body := "您的登录验证码是:" + code +
|
||||
"\n有效期 " + s.cfg.Policy.OTPTTL.String() + ",请勿向他人泄露。\n" +
|
||||
"如未收到邮件,可通过 CLI 子命令 `usernode user otp --username " + u.Username + "` 获取同一验证码。"
|
||||
if err := s.mailer.Send(ctx, u.Email, "登录验证码", body); err != nil {
|
||||
// OTP 邮件失败不阻断登录(PLAN F5);CLI 通道仍可获取同一验证码
|
||||
s.log.Warn("otp mail send failed, cli channel remains available", "username", u.Username, "err", err)
|
||||
}
|
||||
_ = s.audit.Record(ctx, u.ID, u.Username, "user.otp.send", "user", "", nil, ip, model.ResultSuccess)
|
||||
return nil
|
||||
}
|
||||
|
||||
// UserOTPLogin 外部用户 OTP 登录,返回会话 ID。
|
||||
func (s *AuthService) UserOTPLogin(ctx context.Context, username, code, ip, userAgent string) (string, error) {
|
||||
u, err := s.users.GetByUsername(ctx, username)
|
||||
if err != nil {
|
||||
return "", ErrUserUnavailable
|
||||
}
|
||||
if u.Status != model.UserStatusActive {
|
||||
return "", ErrUserUnavailable
|
||||
}
|
||||
ok, err := s.otps.Verify(ctx, u.Username, code)
|
||||
if err != nil {
|
||||
_ = s.audit.Record(ctx, u.ID, u.Username, "user.otp.login", "user", "", map[string]any{"err": err.Error()}, ip, model.ResultFailed)
|
||||
return "", err
|
||||
}
|
||||
if !ok {
|
||||
return "", auth.ErrInvalidCode
|
||||
}
|
||||
now := time.Now()
|
||||
_ = s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", u.ID).Update("last_login_at", &now).Error
|
||||
sid, err := s.createSession(ctx, auth.SessionUserUser, u.ID, ip, userAgent)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
_ = s.audit.Record(ctx, u.ID, u.Username, "user.otp.login", "user", "", nil, ip, model.ResultSuccess)
|
||||
return sid, nil
|
||||
}
|
||||
|
||||
// MeInfo 当前会话对应的主体信息。
|
||||
type MeInfo struct {
|
||||
UserType string `json:"user_type"`
|
||||
ID uint `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Email string `json:"email"`
|
||||
}
|
||||
|
||||
// Me 根据会话 ID 返回当前登录主体。
|
||||
func (s *AuthService) Me(ctx context.Context, sessionID string) (*MeInfo, error) {
|
||||
sess, err := s.sessions.Get(ctx, sessionID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch sess.UserType {
|
||||
case auth.SessionUserAdmin:
|
||||
adm, err := s.admins.GetByID(ctx, sess.RefID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &MeInfo{UserType: sess.UserType, ID: adm.ID, Username: adm.Username, Email: adm.Email}, nil
|
||||
case auth.SessionUserUser:
|
||||
u, err := s.users.GetByID(ctx, sess.RefID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &MeInfo{UserType: sess.UserType, ID: u.ID, Username: u.Username, Email: u.Email}, nil
|
||||
}
|
||||
return nil, ErrSessionInvalid
|
||||
}
|
||||
|
||||
// ErrSessionInvalid 表示会话类型未知。
|
||||
var ErrSessionInvalid = errors.New("service: 会话无效")
|
||||
@@ -0,0 +1,230 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"ws_usernode/internal/auth"
|
||||
)
|
||||
|
||||
// testLogger 丢弃日志输出的测试 logger。
|
||||
func testLogger() *slog.Logger {
|
||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
}
|
||||
|
||||
// recordingMailer 捕获最近一封邮件,便于从重置邮件中提取令牌。
|
||||
type recordingMailer struct {
|
||||
lastTo string
|
||||
lastSubject string
|
||||
lastBody string
|
||||
}
|
||||
|
||||
func (m *recordingMailer) Send(_ context.Context, to, subject, body string) error {
|
||||
m.lastTo = to
|
||||
m.lastSubject = subject
|
||||
m.lastBody = body
|
||||
return nil
|
||||
}
|
||||
|
||||
// newTestAuthService 组装一套完整的认证服务(DB + 内存验证码 + DB OTP/会话/令牌)。
|
||||
func newTestAuthService(t *testing.T) (*AuthService, *fakeSys, *recordingMailer) {
|
||||
t.Helper()
|
||||
db := testDB(t)
|
||||
sys := newFakeSys()
|
||||
userSvc := NewUserService(db, sys, testConfig())
|
||||
adminSvc := NewAdminService(db)
|
||||
auditSvc := NewAuditService(db)
|
||||
cfg := testConfig()
|
||||
|
||||
mailer := &recordingMailer{}
|
||||
otps := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin)
|
||||
captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL)
|
||||
sessions := auth.NewDBSessionStore(db)
|
||||
resets := auth.NewDBResetTokenStore(db)
|
||||
limiter := auth.NewRateLimiter(cfg.Auth.MaxLoginFailures, cfg.Auth.LockDuration)
|
||||
|
||||
svc := NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, testLogger())
|
||||
return svc, sys, mailer
|
||||
}
|
||||
|
||||
func TestAuthServiceAdminLogin(t *testing.T) {
|
||||
svc, _, _ := newTestAuthService(t)
|
||||
ctx := context.Background()
|
||||
mustAdmin(t, NewAdminService(svc.db))
|
||||
|
||||
sid, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test")
|
||||
if err != nil {
|
||||
t.Fatalf("admin login: %v", err)
|
||||
}
|
||||
if sid == "" {
|
||||
t.Fatal("session id should not be empty")
|
||||
}
|
||||
// 错误密码
|
||||
if _, err := svc.AdminLogin(ctx, "root", "wrongpass", "127.0.0.1", "test"); !errors.Is(err, ErrBadCredentials) {
|
||||
t.Fatalf("bad password err = %v, want ErrBadCredentials", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthServiceAdminLoginRateLimited(t *testing.T) {
|
||||
svc, _, _ := newTestAuthService(t)
|
||||
ctx := context.Background()
|
||||
mustAdmin(t, NewAdminService(svc.db))
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
_, _ = svc.AdminLogin(ctx, "root", "wrongpass", "127.0.0.1", "test")
|
||||
}
|
||||
_, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test")
|
||||
if !errors.Is(err, ErrRateLimited) {
|
||||
t.Fatalf("login after lock err = %v, want ErrRateLimited", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthServiceUserOTPFlow(t *testing.T) {
|
||||
svc, sys, mailer := newTestAuthService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// 准备一个外部用户
|
||||
userSvc := NewUserService(svc.db, sys, testConfig())
|
||||
u, err := userSvc.Create(ctx, "zhangsan", "zs@example.com", "", "", 0, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
|
||||
// 获取图形验证码
|
||||
cap, err := svc.captchas.New()
|
||||
if err != nil {
|
||||
t.Fatalf("captcha: %v", err)
|
||||
}
|
||||
|
||||
// 发送 OTP
|
||||
if err := svc.UserOTPSend(ctx, u.Username, cap.ID, cap.Text, "127.0.0.1"); err != nil {
|
||||
t.Fatalf("otp send: %v", err)
|
||||
}
|
||||
if !strings.Contains(mailer.lastBody, "验证码") {
|
||||
t.Fatalf("otp mail body unexpected: %s", mailer.lastBody)
|
||||
}
|
||||
|
||||
// CLI 通道复用同一验证码
|
||||
code, err := svc.otps.Current(ctx, u.Username)
|
||||
if err != nil {
|
||||
t.Fatalf("otp current: %v", err)
|
||||
}
|
||||
if len(code) != 6 {
|
||||
t.Fatalf("otp len = %d, want 6", len(code))
|
||||
}
|
||||
|
||||
// OTP 登录
|
||||
sid, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test")
|
||||
if err != nil {
|
||||
t.Fatalf("otp login: %v", err)
|
||||
}
|
||||
if sid == "" {
|
||||
t.Fatal("session id should not be empty")
|
||||
}
|
||||
|
||||
// 一次性:再次使用失败
|
||||
if _, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test"); err != auth.ErrInvalidCode {
|
||||
t.Fatalf("reuse otp err = %v, want ErrInvalidCode", err)
|
||||
}
|
||||
|
||||
// 验证码错误时 send 被拒绝
|
||||
cap2, _ := svc.captchas.New()
|
||||
if err := svc.UserOTPSend(ctx, u.Username, cap2.ID, "0000", "127.0.0.1"); !errors.Is(err, ErrCaptchaFailed) {
|
||||
t.Fatalf("bad captcha err = %v, want ErrCaptchaFailed", err)
|
||||
}
|
||||
|
||||
// 冷却期内再次 send 被拒
|
||||
cap3, _ := svc.captchas.New()
|
||||
if err := svc.UserOTPSend(ctx, u.Username, cap3.ID, cap3.Text, "127.0.0.1"); !errors.Is(err, auth.ErrCooldown) {
|
||||
t.Fatalf("cooldown err = %v, want ErrCooldown", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthServiceAdminForgotReset(t *testing.T) {
|
||||
svc, _, mailer := newTestAuthService(t)
|
||||
ctx := context.Background()
|
||||
adm := mustAdmin(t, NewAdminService(svc.db))
|
||||
|
||||
// forgot → 邮件应包含重置链接
|
||||
if err := svc.AdminForgot(ctx, adm.Username, "127.0.0.1"); err != nil {
|
||||
t.Fatalf("forgot: %v", err)
|
||||
}
|
||||
if mailer.lastTo != adm.Email {
|
||||
t.Fatalf("mail to = %q, want %q", mailer.lastTo, adm.Email)
|
||||
}
|
||||
// 从邮件 body 提取 token
|
||||
idx := strings.Index(mailer.lastBody, "token=")
|
||||
if idx < 0 {
|
||||
t.Fatalf("reset link missing token: %s", mailer.lastBody)
|
||||
}
|
||||
token := mailer.lastBody[idx+len("token="):]
|
||||
token = strings.TrimSpace(strings.SplitN(token, "\n", 2)[0])
|
||||
|
||||
// reset
|
||||
if err := svc.AdminReset(ctx, token, "NewPassw0rd", "127.0.0.1"); err != nil {
|
||||
t.Fatalf("reset: %v", err)
|
||||
}
|
||||
// 新密码可登录
|
||||
adminSvc := NewAdminService(svc.db)
|
||||
if _, err := adminSvc.Login(ctx, adm.Username, "NewPassw0rd"); err != nil {
|
||||
t.Fatalf("login with new password: %v", err)
|
||||
}
|
||||
// 旧密码失效
|
||||
if _, err := adminSvc.Login(ctx, adm.Username, "Passw0rd"); !errors.Is(err, ErrBadCredentials) {
|
||||
t.Fatalf("login with old password err = %v, want ErrBadCredentials", err)
|
||||
}
|
||||
// 令牌一次性
|
||||
if err := svc.AdminReset(ctx, token, "AgainPassw0rd", "127.0.0.1"); !errors.Is(err, auth.ErrResetTokenInvalid) {
|
||||
t.Fatalf("reuse token err = %v, want ErrResetTokenInvalid", err)
|
||||
}
|
||||
// 不存在的用户 forgot 也应成功(防枚举)
|
||||
if err := svc.AdminForgot(ctx, "ghost", "127.0.0.1"); err != nil {
|
||||
t.Fatalf("forgot ghost: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthServiceMeLogout(t *testing.T) {
|
||||
svc, sys, _ := newTestAuthService(t)
|
||||
ctx := context.Background()
|
||||
mustAdmin(t, NewAdminService(svc.db))
|
||||
|
||||
sid, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test")
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
info, err := svc.Me(ctx, sid)
|
||||
if err != nil {
|
||||
t.Fatalf("me: %v", err)
|
||||
}
|
||||
if info.UserType != "admin" || info.Username != "root" {
|
||||
t.Fatalf("me info = %+v", info)
|
||||
}
|
||||
if err := svc.Logout(ctx, sid); err != nil {
|
||||
t.Fatalf("logout: %v", err)
|
||||
}
|
||||
if _, err := svc.Me(ctx, sid); !errors.Is(err, auth.ErrSessionNotFound) {
|
||||
t.Fatalf("me after logout err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
|
||||
// 外部用户 me
|
||||
userSvc := NewUserService(svc.db, sys, testConfig())
|
||||
u, _ := userSvc.Create(ctx, "lisi", "ls@example.com", "", "", 0, 0)
|
||||
cap, _ := svc.captchas.New()
|
||||
_ = svc.UserOTPSend(ctx, u.Username, cap.ID, cap.Text, "127.0.0.1")
|
||||
code, _ := svc.otps.Current(ctx, u.Username)
|
||||
usid, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test")
|
||||
if err != nil {
|
||||
t.Fatalf("user login: %v", err)
|
||||
}
|
||||
info, err = svc.Me(ctx, usid)
|
||||
if err != nil {
|
||||
t.Fatalf("user me: %v", err)
|
||||
}
|
||||
if info.UserType != "user" || info.Username != "ext_lisi" {
|
||||
t.Fatalf("user me info = %+v", info)
|
||||
}
|
||||
}
|
||||
+229
-17
@@ -4,9 +4,11 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"ws_usernode/internal/config"
|
||||
"ws_usernode/internal/model"
|
||||
"ws_usernode/internal/pkg"
|
||||
"ws_usernode/internal/system"
|
||||
@@ -14,29 +16,30 @@ import (
|
||||
|
||||
// 用户服务错误。
|
||||
var (
|
||||
ErrUserNotFound = errors.New("service: 用户不存在")
|
||||
ErrUserExists = errors.New("service: 用户名已存在")
|
||||
ErrUserNotFound = errors.New("service: 用户不存在")
|
||||
ErrUserExists = errors.New("service: 用户名已存在")
|
||||
ErrUserExpired = errors.New("service: 用户已过期,请先延期")
|
||||
ErrUserDisabled = errors.New("service: 用户已禁用,无法操作")
|
||||
ErrSystemAccountMissing = errors.New("service: 系统账号不存在,无法操作")
|
||||
)
|
||||
|
||||
// UserService 外部用户生命周期服务。
|
||||
// M0 提供查询与创建骨架;建号(useradd)、禁用/延期等系统操作 M1 接入 system.Manager。
|
||||
// UserService 外部用户生命周期服务:DB 记录 + system.Manager 系统账号操作。
|
||||
type UserService struct {
|
||||
db *gorm.DB
|
||||
sys system.Manager
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
// NewUserService 创建用户服务。
|
||||
func NewUserService(db *gorm.DB, sys system.Manager) *UserService {
|
||||
return &UserService{db: db, sys: sys}
|
||||
func NewUserService(db *gorm.DB, sys system.Manager, cfg *config.Config) *UserService {
|
||||
return &UserService{db: db, sys: sys, cfg: cfg}
|
||||
}
|
||||
|
||||
// GetByUsername 按用户名查询外部用户(含或不含 ext_ 前缀均可)。
|
||||
func (s *UserService) GetByUsername(ctx context.Context, username string) (*model.User, error) {
|
||||
if !strings.HasPrefix(username, "ext_") {
|
||||
username = "ext_" + username
|
||||
}
|
||||
full := normalizeName(username, s.cfg.System.UserPrefix)
|
||||
var u model.User
|
||||
if err := s.db.WithContext(ctx).First(&u, "username = ?", username).Error; err != nil {
|
||||
if err := s.db.WithContext(ctx).First(&u, "username = ?", full).Error; err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
@@ -45,16 +48,68 @@ func (s *UserService) GetByUsername(ctx context.Context, username string) (*mode
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
// Create 创建外部用户记录并调用系统层建号。M0 阶段系统层为 dry-run。
|
||||
// username 为不含前缀的申请名,内部加 ext_ 前缀。
|
||||
func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttlSeconds int64) (*model.User, error) {
|
||||
// GetByID 按 ID 查询外部用户。
|
||||
func (s *UserService) GetByID(ctx context.Context, id uint) (*model.User, error) {
|
||||
var u model.User
|
||||
if err := s.db.WithContext(ctx).First(&u, "id = ?", id).Error; err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
// UserFilter 用户列表筛选条件。
|
||||
type UserFilter struct {
|
||||
Status string // active / disabled / expired,空为全部
|
||||
Supervisor string // 挂靠老师模糊匹配
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// List 分页查询用户(admin)。
|
||||
func (s *UserService) List(ctx context.Context, f UserFilter) ([]model.User, int64, error) {
|
||||
q := s.db.WithContext(ctx).Model(&model.User{})
|
||||
if f.Status != "" {
|
||||
q = q.Where("status = ?", f.Status)
|
||||
}
|
||||
if f.Supervisor != "" {
|
||||
q = q.Where("supervisor LIKE ?", "%"+f.Supervisor+"%")
|
||||
}
|
||||
var total int64
|
||||
if err := q.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
page, size := f.Page, f.PageSize
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if size < 1 {
|
||||
size = 20
|
||||
}
|
||||
if size > 100 {
|
||||
size = 100
|
||||
}
|
||||
var users []model.User
|
||||
if err := q.Order("id DESC").Offset((page - 1) * size).Limit(size).Find(&users).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return users, total, nil
|
||||
}
|
||||
|
||||
// Create 创建外部用户:DB 记录 + 系统账号(useradd + passwd -l)。
|
||||
// username 不含前缀;ttl 为有效期时长,0 表示用配置默认(90 天)。
|
||||
// 系统建号失败时回滚 DB 记录,保证两侧一致。
|
||||
func (s *UserService) Create(ctx context.Context, username, email, supervisor, purpose string, ttl time.Duration, createdBy uint) (*model.User, error) {
|
||||
if err := pkg.ValidateUserName(username); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := pkg.ValidateEmail(email); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
full := "ext_" + username
|
||||
username = strings.TrimSpace(username)
|
||||
full := s.cfg.System.UserPrefix + username
|
||||
var count int64
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", full).Count(&count).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -62,21 +117,178 @@ func (s *UserService) Create(ctx context.Context, username, email, supervisor, p
|
||||
if count > 0 {
|
||||
return nil, ErrUserExists
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = s.cfg.Policy.DefaultTTL
|
||||
}
|
||||
expireAt := time.Now().Add(ttl)
|
||||
u := &model.User{
|
||||
Username: full,
|
||||
Email: email,
|
||||
Supervisor: supervisor,
|
||||
Purpose: purpose,
|
||||
Status: model.UserStatusActive,
|
||||
Shell: "/bin/sh",
|
||||
ExpireAt: &expireAt,
|
||||
Shell: s.cfg.System.Shell,
|
||||
CreatedBy: createdBy,
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Create(u).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 系统账号创建(dry-run / 真实),失败时回滚 DB 记录
|
||||
if err := s.sys.CreateUser(ctx, system.Account{Username: full}); err != nil {
|
||||
// 系统账号创建(dry-run / 直接 / sudo),失败时回滚 DB 记录
|
||||
if err := s.sys.CreateUser(ctx, system.Account{Username: full, Shell: s.cfg.System.Shell}); err != nil {
|
||||
_ = s.db.WithContext(ctx).Delete(u).Error
|
||||
return nil, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// Update 更新外部用户信息(仅更新非 nil 字段;邮箱由管理员修改,用户不可自助改)。
|
||||
func (s *UserService) Update(ctx context.Context, id uint, email, supervisor, purpose *string) (*model.User, error) {
|
||||
if _, err := s.GetByID(ctx, id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
updates := make(map[string]any)
|
||||
if email != nil {
|
||||
if err := pkg.ValidateEmail(*email); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
updates["email"] = *email
|
||||
}
|
||||
if supervisor != nil {
|
||||
updates["supervisor"] = *supervisor
|
||||
}
|
||||
if purpose != nil {
|
||||
updates["purpose"] = *purpose
|
||||
}
|
||||
if len(updates) > 0 {
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return s.GetByID(ctx, id)
|
||||
}
|
||||
|
||||
// systemAccountOK 检查系统账号存在;dry-run 模式跳过检查(演练流程)。
|
||||
func (s *UserService) systemAccountOK(ctx context.Context, username string) bool {
|
||||
if s.cfg.System.DryRun {
|
||||
return true
|
||||
}
|
||||
ok, err := s.sys.Exists(ctx, username)
|
||||
return err == nil && ok
|
||||
}
|
||||
|
||||
// Disable 禁用用户:DB 置 disabled + 清空 authorized_keys(SSH 立即失效)。
|
||||
func (s *UserService) Disable(ctx context.Context, id uint) error {
|
||||
u, err := s.GetByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if u.Status == model.UserStatusDisabled {
|
||||
return nil // 幂等
|
||||
}
|
||||
if !s.systemAccountOK(ctx, u.Username) {
|
||||
return ErrSystemAccountMissing
|
||||
}
|
||||
if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Update("status", model.UserStatusDisabled).Error
|
||||
}
|
||||
|
||||
// Enable 启用用户:DB 置 active + 按 DB 状态重写 authorized_keys。
|
||||
// 已过期的用户需先延期(Extend)。
|
||||
func (s *UserService) Enable(ctx context.Context, id uint) error {
|
||||
u, err := s.GetByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if u.Status == model.UserStatusActive {
|
||||
return nil // 幂等
|
||||
}
|
||||
if u.ExpireAt != nil && time.Now().After(*u.ExpireAt) {
|
||||
return ErrUserExpired
|
||||
}
|
||||
if !s.systemAccountOK(ctx, u.Username) {
|
||||
return ErrSystemAccountMissing
|
||||
}
|
||||
// 恢复有效密钥(M1 阶段用户尚无密钥,M2 接入后按 DB 同步)
|
||||
keys, err := s.activeKeys(ctx, u.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, keys); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Update("status", model.UserStatusActive).Error
|
||||
}
|
||||
|
||||
// Extend 延长有效期:重设 expire_at(days<=0 用配置默认 TTL)。
|
||||
// 已过期用户在回收期内可经此恢复(PLAN §2.2),恢复后同步密钥。
|
||||
func (s *UserService) Extend(ctx context.Context, id uint, days int) error {
|
||||
u, err := s.GetByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ttl := time.Duration(days) * 24 * time.Hour
|
||||
if days <= 0 {
|
||||
ttl = s.cfg.Policy.DefaultTTL
|
||||
}
|
||||
newExpire := time.Now().Add(ttl)
|
||||
updates := map[string]any{"expire_at": newExpire}
|
||||
if u.Status == model.UserStatusExpired {
|
||||
if !s.systemAccountOK(ctx, u.Username) {
|
||||
return ErrSystemAccountMissing
|
||||
}
|
||||
keys, err := s.activeKeys(ctx, u.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, keys); err != nil {
|
||||
return err
|
||||
}
|
||||
updates["status"] = model.UserStatusActive
|
||||
}
|
||||
return s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error
|
||||
}
|
||||
|
||||
// Delete 删除并回收用户:删除系统账号(userdel -r)+ 家目录 + 密钥记录,
|
||||
// 保留审计。系统账号已不存在时仍完成 DB 清理。
|
||||
func (s *UserService) Delete(ctx context.Context, id uint) error {
|
||||
u, err := s.GetByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if s.systemAccountOK(ctx, u.Username) {
|
||||
if err := s.sys.RemoveUser(ctx, u.Username); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("user_id = ?", u.ID).Delete(&model.SSHKey{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.User{}, "id = ?", u.ID).Error
|
||||
})
|
||||
}
|
||||
|
||||
// activeKeys 返回用户当前有效(active)密钥,供授权同步(M2 完善密钥管理)。
|
||||
func (s *UserService) activeKeys(ctx context.Context, userID uint) ([]system.Key, error) {
|
||||
var rows []model.SSHKey
|
||||
if err := s.db.WithContext(ctx).Where("user_id = ? AND status = ?", userID, model.StatusActive).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keys := make([]system.Key, 0, len(rows))
|
||||
for _, k := range rows {
|
||||
keys = append(keys, system.Key{Type: k.KeyType, PublicKey: k.PublicKey})
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// normalizeName 补全系统账号前缀(如 ext_)。
|
||||
func normalizeName(username, prefix string) string {
|
||||
name := strings.TrimSpace(username)
|
||||
if !strings.HasPrefix(name, prefix) {
|
||||
return prefix + name
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user