diff --git a/Makefile b/Makefile index bddd808..71d008d 100644 --- a/Makefile +++ b/Makefile @@ -13,7 +13,7 @@ NET_HOST := --network=host GO ?= go PODMAN ?= podman BIN := bin/usernode -VERSION ?= 0.2.0-m1 +VERSION ?= 0.3.0-m2 LDFLAGS := -s -w -X main.version=$(VERSION) GOFLAGS := -trimpath diff --git a/cmd/usernode/app.go b/cmd/usernode/app.go index beda986..eab1db9 100644 --- a/cmd/usernode/app.go +++ b/cmd/usernode/app.go @@ -82,9 +82,10 @@ func cmdServe(args []string) error { return err } - sys := system.New(cfg.System) + sys := system.New(cfg.System, log) adminSvc := service.NewAdminService(db) userSvc := service.NewUserService(db, sys, cfg) + keySvc := service.NewKeyService(db, sys, cfg) auditSvc := service.NewAuditService(db) // 认证依赖:DB OTP/会话/重置令牌存储 + 内存图形验证码/登录限速器 + 邮件 @@ -108,7 +109,7 @@ func cmdServe(args []string) error { sched.Start() defer sched.Stop() - h := api.New(cfg, authSvc, userSvc, auditSvc) + h := api.New(cfg, authSvc, userSvc, keySvc, auditSvc) r := router.New(cfg, h, sessions, log) srv := server.New(cfg.Server.Listen, r, log) diff --git a/cmd/usernode/user.go b/cmd/usernode/user.go index 6057177..8c40c7b 100644 --- a/cmd/usernode/user.go +++ b/cmd/usernode/user.go @@ -48,7 +48,7 @@ func userOTP(args []string) error { } // 校验用户存在(不存在时返回友好错误,避免暴露账号是否存在的枚举) - if _, err := service.NewUserService(db, system.New(cfg.System), cfg).GetByUsername(context.Background(), name); err != nil { + if _, err := service.NewUserService(db, system.New(cfg.System, log), cfg).GetByUsername(context.Background(), name); err != nil { return fmt.Errorf("用户不存在或不可用: %w", err) } diff --git a/config.example.toml b/config.example.toml index 0de141b..f64a12c 100644 --- a/config.example.toml +++ b/config.example.toml @@ -44,7 +44,7 @@ password = "" from = "usernode@example.com" [system] -sudo = false # 生产 true:经 sudo -n 执行 useradd/usermod/userdel/passwd(需 deploy/sudoers) +sudo = false # 生产 true:经 sudo -n 执行白名单命令(账号生命周期 + authorized_keys 同步,需 deploy/sudoers) dry_run = true # 开发演练 true:只打印计划命令不执行;false 且 sudo=false 时直接执行(容器/测试用户验证) user_prefix = "ext_" # 外部用户系统账号统一前缀 group = "external" # 外部用户统一组 diff --git a/deploy/sudoers.example b/deploy/sudoers.example index 417f6a3..ac22c75 100644 --- a/deploy/sudoers.example +++ b/deploy/sudoers.example @@ -4,17 +4,21 @@ # 节点进程以 usernode 用户运行,仅允许以 root 执行固定命令(禁任意 shell), # 命令参数由程序内强校验(pkg.ValidateSystemAccount 等),见 PLAN §9。 # -# 注意:以下命令路径基于 Debian/Ubuntu(/usr/sbin)。Alpine 为 /usr/sbin; -# 请按发行版调整,并确保 usernode 用户无 NOPASSWD 的通用提权入口。 +# 注意:以下命令路径基于 Debian/Ubuntu(/usr/sbin、/usr/bin、/bin)。 +# Alpine 为 /usr/sbin、/usr/bin、/bin;请按发行版调整,并确保 usernode +# 用户无 NOPASSWD 的通用提权入口。 usernode ALL=(root) NOPASSWD: /usr/sbin/useradd, /usr/sbin/usermod, \ - /usr/sbin/userdel, /usr/bin/passwd + /usr/sbin/userdel, /usr/bin/passwd, \ + /bin/mkdir, /bin/chmod, /bin/chown, /usr/bin/install # 说明: -# - useradd -m -d -s -g external 创建账号 +# - useradd -m -d -s -g external 创建账号 # - usermod 预留(如 usermod -e 过期),M4 回收期使用 -# - userdel -r 删除账号及家目录 -# - passwd -l / -u 锁定/解锁口令 +# - userdel -r 删除账号及家目录 +# - passwd -l / -u 锁定/解锁口令 +# - mkdir -p / chown / chmod:创建并修正 ~/.ssh 目录属主(sshd StrictModes) +# - install -o -g external -m 0600 :authorized_keys 原子落位 # - 不授予 chsh/其他命令的任意执行;若需变更默认 shell 请收紧为固定参数 # # 生产禁止 system.sudo=false 的 direct 模式:必须显式配置 diff --git a/internal/api/api_test.go b/internal/api/api_test.go index cb06930..4625e7f 100644 --- a/internal/api/api_test.go +++ b/internal/api/api_test.go @@ -3,15 +3,20 @@ package api_test import ( "bytes" "context" + "crypto/ed25519" + "crypto/rand" "encoding/json" "io" "log/slog" "net/http" "net/http/httptest" "strconv" + "strings" + "sync" "testing" "github.com/gin-gonic/gin" + "golang.org/x/crypto/ssh" "gorm.io/gorm" "ws_usernode/internal/api" @@ -37,7 +42,32 @@ func (m *recordingMailer) Send(_ context.Context, to, subject, body string) erro return nil } -// testApp 完整组装的应用(SQLite 内存库 + dry-run 系统层)。 +// recordingSys 记录 authorized_keys 同步内容,用于断言上传/吊销/禁用即时生效。 +type recordingSys struct { + mu sync.Mutex + keys map[string][]system.Key // username -> 最后一次同步的密钥 +} + +func newRecordingSys() *recordingSys { return &recordingSys{keys: map[string][]system.Key{}} } + +func (s *recordingSys) CreateUser(_ context.Context, _ system.Account) error { return nil } +func (s *recordingSys) RemoveUser(_ context.Context, _ string) error { return nil } +func (s *recordingSys) SetLock(_ context.Context, _ string, _ bool) error { return nil } +func (s *recordingSys) Exists(_ context.Context, _ string) (bool, error) { return true, nil } +func (s *recordingSys) SyncAuthorizedKeys(_ context.Context, username string, keys []system.Key) error { + s.mu.Lock() + defer s.mu.Unlock() + s.keys[username] = keys + return nil +} + +func (s *recordingSys) synced(username string) []system.Key { + s.mu.Lock() + defer s.mu.Unlock() + return s.keys[username] +} + +// testApp 完整组装的应用(SQLite 内存库 + 可注入系统层)。 type testApp struct { r http.Handler db *gorm.DB @@ -47,6 +77,11 @@ type testApp struct { } func setupTestApp(t *testing.T) *testApp { + return setupTestAppWithSys(t, system.New(config.Default().System, slog.New(slog.NewTextHandler(io.Discard, nil)))) +} + +// setupTestAppWithSys 允许注入 system.Manager(观察 authorized_keys 同步等)。 +func setupTestAppWithSys(t *testing.T, sys system.Manager) *testApp { t.Helper() gin.SetMode(gin.TestMode) db, err := model.Open("sqlite", ":memory:", false) @@ -59,9 +94,9 @@ func setupTestApp(t *testing.T) *testApp { cfg := config.Default() cfg.System.DryRun = true // 集成测试走 dry-run,不触碰真实系统账号 - sys := system.New(cfg.System) adminSvc := service.NewAdminService(db) userSvc := service.NewUserService(db, sys, cfg) + keySvc := service.NewKeyService(db, sys, cfg) auditSvc := service.NewAuditService(db) captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL) @@ -73,7 +108,7 @@ func setupTestApp(t *testing.T) *testApp { 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) + h := api.New(cfg, authSvc, userSvc, keySvc, auditSvc) r := router.New(cfg, h, sessions, log) if _, err := adminSvc.Create(context.Background(), "root", "Passw0rd", "root@example.com"); err != nil { @@ -324,3 +359,124 @@ func TestAPIAdminForgotReset(t *testing.T) { func itoa(u uint) string { return strconv.FormatUint(uint64(u), 10) } + +func TestAPIKeyLifecycle(t *testing.T) { + rec := newRecordingSys() + app := setupTestAppWithSys(t, rec) + + // 管理员创建外部用户 + 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": "wangwu", "email": "ww@example.com"}, ck) + if w.Code != http.StatusOK { + t.Fatalf("create user status = %d, body=%s", w.Code, w.Body.String()) + } + userID := uint(decodeBody(t, w)["data"].(map[string]any)["id"].(float64)) + + // 外部用户 OTP 登录 + cap, err := app.captchas.New() + if err != nil { + t.Fatalf("captcha new: %v", err) + } + w = app.doJSON(http.MethodPost, "/api/v1/auth/otp/send", map[string]any{"username": "ext_wangwu", "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()) + } + code, err := app.otps.Current(context.Background(), "ext_wangwu") + if err != nil { + t.Fatalf("otp current: %v", err) + } + w = app.doJSON(http.MethodPost, "/api/v1/auth/otp/login", map[string]string{"username": "ext_wangwu", "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/keys → 401 + w = app.doJSON(http.MethodGet, "/api/v1/me/keys", nil) + if w.Code != http.StatusUnauthorized { + t.Fatalf("unauth me/keys status = %d, want 401", w.Code) + } + + // 上传公钥 → 同步到 authorized_keys + pub := testSSHPubKey(t) + w = app.doJSON(http.MethodPost, "/api/v1/me/keys", map[string]any{"name": "workstation", "public_key": pub}, userCk) + if w.Code != http.StatusOK { + t.Fatalf("create key status = %d, body=%s", w.Code, w.Body.String()) + } + key := decodeBody(t, w)["data"].(map[string]any) + keyID := uint(key["id"].(float64)) + if key["fingerprint"] == "" || key["status"] != "active" { + t.Fatalf("key fields: %v", key) + } + if synced := rec.synced("ext_wangwu"); len(synced) != 1 { + t.Fatalf("after create synced = %+v, want 1 key", synced) + } + + // 重复上传同一公钥 → 409 + w = app.doJSON(http.MethodPost, "/api/v1/me/keys", map[string]any{"name": "dup", "public_key": pub}, userCk) + if w.Code != http.StatusConflict { + t.Fatalf("duplicate key status = %d, want 409", w.Code) + } + + // 列表 + w = app.doJSON(http.MethodGet, "/api/v1/me/keys", nil, userCk) + if w.Code != http.StatusOK { + t.Fatalf("list keys status = %d", w.Code) + } + if items := decodeBody(t, w)["data"].(map[string]any)["items"].([]any); len(items) != 1 { + t.Fatalf("list items = %d, want 1", len(items)) + } + + // 重命名 + w = app.doJSON(http.MethodPatch, "/api/v1/me/keys/"+itoa(keyID), map[string]any{"name": "home-laptop"}, userCk) + if w.Code != http.StatusOK { + t.Fatalf("rename status = %d, body=%s", w.Code, w.Body.String()) + } + if name := decodeBody(t, w)["data"].(map[string]any)["name"]; name != "home-laptop" { + t.Fatalf("renamed = %v", name) + } + + // 管理员查看用户密钥;外部用户访问 admin 接口 → 403 + w = app.doJSON(http.MethodGet, "/api/v1/users/"+itoa(userID)+"/keys", nil, ck) + if w.Code != http.StatusOK { + t.Fatalf("admin list keys status = %d, body=%s", w.Code, w.Body.String()) + } + if items := decodeBody(t, w)["data"].(map[string]any)["items"].([]any); len(items) != 1 { + t.Fatalf("admin items = %d, want 1", len(items)) + } + w = app.doJSON(http.MethodGet, "/api/v1/users/"+itoa(userID)+"/keys", nil, userCk) + if w.Code != http.StatusForbidden { + t.Fatalf("user access admin keys status = %d, want 403", w.Code) + } + + // 吊销 → authorized_keys 清空(立即失效),状态 revoked;再次吊销幂等 + w = app.doJSON(http.MethodDelete, "/api/v1/me/keys/"+itoa(keyID), nil, userCk) + if w.Code != http.StatusOK { + t.Fatalf("revoke status = %d, body=%s", w.Code, w.Body.String()) + } + if status := decodeBody(t, w)["data"].(map[string]any)["status"]; status != "revoked" { + t.Fatalf("revoked status = %v", status) + } + if synced := rec.synced("ext_wangwu"); len(synced) != 0 { + t.Fatalf("after revoke synced = %+v, want empty", synced) + } + w = app.doJSON(http.MethodDelete, "/api/v1/me/keys/"+itoa(keyID), nil, userCk) + if w.Code != http.StatusOK { + t.Fatalf("revoke again status = %d", w.Code) + } +} + +// testSSHPubKey 生成一条合法的 ed25519 公钥行。 +func testSSHPubKey(t *testing.T) string { + t.Helper() + pub, _, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatalf("gen ed25519: %v", err) + } + sshPub, err := ssh.NewPublicKey(pub) + if err != nil { + t.Fatalf("ssh key: %v", err) + } + return strings.TrimSpace(string(ssh.MarshalAuthorizedKey(sshPub))) +} diff --git a/internal/api/handler.go b/internal/api/handler.go index 3dc7bca..818bfa5 100644 --- a/internal/api/handler.go +++ b/internal/api/handler.go @@ -1,5 +1,6 @@ // Package api 为 HTTP handler 层(RESTful v1)。 -// M1 覆盖认证(管理员/外部用户登录、会话、OTP)与用户管理 CRUD。 +// M1 覆盖认证(管理员/外部用户登录、会话、OTP)与用户管理 CRUD; +// M2 覆盖 SSH 公钥管理(上传/重命名/吊销/列表)。 package api import ( @@ -20,13 +21,14 @@ type Handler struct { Health *HealthHandler Auth *AuthHandler User *UserHandler + Key *KeyHandler - authSvc *service.AuthService + authSvc *service.AuthService auditSvc *service.AuditService } // New 创建 handler 集合。 -func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, auditSvc *service.AuditService) *Handler { +func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.UserService, keySvc *service.KeyService, auditSvc *service.AuditService) *Handler { h := &Handler{ Health: &HealthHandler{startedAt: time.Now()}, authSvc: authSvc, @@ -34,6 +36,7 @@ func New(cfg *config.Config, authSvc *service.AuthService, userSvc *service.User } h.Auth = &AuthHandler{svc: authSvc, cfg: cfg} h.User = &UserHandler{svc: userSvc, cfg: cfg, h: h} + h.Key = &KeyHandler{svc: keySvc, cfg: cfg, h: h} return h } @@ -60,7 +63,7 @@ type HealthHandler struct { func (h *HealthHandler) Healthz(c *gin.Context) { c.JSON(http.StatusOK, gin.H{ "status": "ok", - "version": "0.2.0-m1", + "version": "0.3.0-m2", "uptime": time.Since(h.startedAt).String(), "go": runtime.Version(), "timestamp": time.Now().UTC().Format(time.RFC3339), diff --git a/internal/api/key.go b/internal/api/key.go new file mode 100644 index 0000000..05caffb --- /dev/null +++ b/internal/api/key.go @@ -0,0 +1,150 @@ +package api + +import ( + "errors" + "net/http" + "strconv" + + "github.com/gin-gonic/gin" + + "ws_usernode/internal/config" + "ws_usernode/internal/model" + "ws_usernode/internal/service" +) + +// KeyHandler SSH 公钥接口:外部用户自助管理(/me/keys)+ 管理员查看(/users/:id/keys)。 +// 密钥仅用户上传(管理员不代签);吊销后立即从 authorized_keys 移除。 +type KeyHandler struct { + svc *service.KeyService + cfg *config.Config + h *Handler // 访问审计 helper +} + +// KeyCreateRequest 上传公钥。 +type KeyCreateRequest struct { + Name string `json:"name" binding:"required"` // 显示名称,1~64 字符 + PublicKey string `json:"public_key" binding:"required"` +} + +// KeyRenameRequest 重命名。 +type KeyRenameRequest struct { + Name string `json:"name" binding:"required"` +} + +// currentUserID 从会话取当前外部用户 ID(/me/keys 均为 user 会话)。 +func (h *KeyHandler) currentUserID(c *gin.Context) uint { + if sess := sessionFrom(c); sess != nil { + return sess.RefID + } + return 0 +} + +// ListMine GET /me/keys —— 我的密钥列表。 +func (h *KeyHandler) ListMine(c *gin.Context) { + keys, err := h.svc.ListByUser(c.Request.Context(), h.currentUserID(c)) + if err != nil { + fail(c, http.StatusInternalServerError, err.Error()) + return + } + ok(c, gin.H{"items": keys}) +} + +// Create POST /me/keys —— 上传公钥(类型/长度/重复校验 + 同步 authorized_keys)。 +func (h *KeyHandler) Create(c *gin.Context) { + var req KeyCreateRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + uid := h.currentUserID(c) + k, err := h.svc.Create(c.Request.Context(), uid, req.Name, req.PublicKey, uid) + if err != nil { + h.h.audit(c, "key.create", "ssh_key", "", map[string]any{"name": req.Name, "err": err.Error()}, model.ResultFailed) + switch { + case errors.Is(err, service.ErrKeyInvalid): + fail(c, http.StatusBadRequest, err.Error()) + case errors.Is(err, service.ErrKeyDuplicate): + fail(c, http.StatusConflict, err.Error()) + case errors.Is(err, service.ErrUserNotFound): + fail(c, http.StatusNotFound, err.Error()) + case errors.Is(err, service.ErrUserNotActive): + fail(c, http.StatusConflict, err.Error()) + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + h.h.audit(c, "key.create", "ssh_key", strconv.FormatUint(uint64(k.ID), 10), map[string]any{"name": k.Name, "fingerprint": k.Fingerprint}, model.ResultSuccess) + ok(c, k) +} + +// Rename PATCH /me/keys/:id —— 重命名(仅元数据,不影响 authorized_keys)。 +func (h *KeyHandler) Rename(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的密钥 ID") + return + } + var req KeyRenameRequest + if err := c.ShouldBindJSON(&req); err != nil { + fail(c, http.StatusBadRequest, "请求参数不合法: "+err.Error()) + return + } + k, err := h.svc.Rename(c.Request.Context(), uint(id), h.currentUserID(c), req.Name) + if err != nil { + h.h.audit(c, "key.rename", "ssh_key", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed) + switch { + case errors.Is(err, service.ErrKeyNotFound): + fail(c, http.StatusNotFound, err.Error()) + case errors.Is(err, service.ErrKeyInvalid): + fail(c, http.StatusBadRequest, err.Error()) + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + h.h.audit(c, "key.rename", "ssh_key", c.Param("id"), map[string]any{"name": k.Name}, model.ResultSuccess) + ok(c, k) +} + +// Revoke DELETE /me/keys/:id —— 吊销(从 authorized_keys 移除,立即失效)。 +func (h *KeyHandler) Revoke(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的密钥 ID") + return + } + k, err := h.svc.Revoke(c.Request.Context(), uint(id), h.currentUserID(c)) + if err != nil { + h.h.audit(c, "key.revoke", "ssh_key", c.Param("id"), map[string]any{"err": err.Error()}, model.ResultFailed) + switch { + case errors.Is(err, service.ErrKeyNotFound): + fail(c, http.StatusNotFound, err.Error()) + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + h.h.audit(c, "key.revoke", "ssh_key", c.Param("id"), map[string]any{"fingerprint": k.Fingerprint}, model.ResultSuccess) + ok(c, k) +} + +// ListForUser GET /users/:id/keys —— 管理员查看用户密钥。 +func (h *KeyHandler) ListForUser(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + fail(c, http.StatusBadRequest, "无效的用户 ID") + return + } + keys, err := h.svc.ListByUser(c.Request.Context(), uint(id)) + if err != nil { + switch { + case errors.Is(err, service.ErrUserNotFound): + fail(c, http.StatusNotFound, err.Error()) + default: + fail(c, http.StatusInternalServerError, err.Error()) + } + return + } + ok(c, gin.H{"items": keys}) +} diff --git a/internal/router/router.go b/internal/router/router.go index 5547021..d7efdaa 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -54,6 +54,16 @@ func New(cfg *config.Config, h *api.Handler, sessions auth.SessionStore, log *sl users.POST("/:id/enable", h.User.Enable) users.POST("/:id/extend", h.User.Extend) users.DELETE("/:id", h.User.Delete) + users.GET("/:id/keys", h.Key.ListForUser) + } + + // 我的密钥(外部用户,自助管理;仅用户上传,管理员不代签) + me := v1.Group("/me", sessionMiddleware(sessions), requireUserType(auth.SessionUserUser)) + { + me.GET("/keys", h.Key.ListMine) + me.POST("/keys", h.Key.Create) + me.PATCH("/keys/:id", h.Key.Rename) + me.DELETE("/keys/:id", h.Key.Revoke) } // 后续里程碑 diff --git a/internal/service/key.go b/internal/service/key.go new file mode 100644 index 0000000..e283b57 --- /dev/null +++ b/internal/service/key.go @@ -0,0 +1,219 @@ +package service + +import ( + "context" + "crypto/rsa" + "encoding/base64" + "errors" + "fmt" + "strings" + "time" + + "golang.org/x/crypto/ssh" + "gorm.io/gorm" + + "ws_usernode/internal/config" + "ws_usernode/internal/model" + "ws_usernode/internal/system" +) + +// 密钥服务错误。 +var ( + ErrKeyNotFound = errors.New("service: 密钥不存在") + ErrKeyDuplicate = errors.New("service: 该公钥已存在") + ErrKeyInvalid = errors.New("service: 公钥格式不合法") + ErrUserNotActive = errors.New("service: 用户未处于可用状态,无法添加密钥") +) + +// maxPublicKeyLen 公钥输入上限(正常公钥约 100~700 字节,防止异常大输入)。 +const maxPublicKeyLen = 8192 + +// KeyService SSH 公钥管理。密钥仅用户自行上传(管理员不代签,PLAN §2.3); +// 每次变更后以 DB 状态全量重写 authorized_keys(system 层原子写 + 并发锁), +// 吊销密钥即从文件移除、立即失效(PLAN F3)。 +type KeyService struct { + db *gorm.DB + sys system.Manager + cfg *config.Config +} + +// NewKeyService 创建密钥服务。 +func NewKeyService(db *gorm.DB, sys system.Manager, cfg *config.Config) *KeyService { + return &KeyService{db: db, sys: sys, cfg: cfg} +} + +// user 查询外部用户记录。 +func (s *KeyService) user(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 +} + +// ListByUser 返回用户全部密钥(最新在前)。用户不存在时报 ErrUserNotFound。 +func (s *KeyService) ListByUser(ctx context.Context, userID uint) ([]model.SSHKey, error) { + if _, err := s.user(ctx, userID); err != nil { + return nil, err + } + var keys []model.SSHKey + if err := s.db.WithContext(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&keys).Error; err != nil { + return nil, err + } + return keys, nil +} + +// Create 上传公钥:校验(类型/长度/重复)→ 落 DB → 全量同步 authorized_keys。 +// 同步失败时回滚 DB 记录,保证两侧一致(同 user.Create 模式)。 +func (s *KeyService) Create(ctx context.Context, userID uint, name, publicKey string, createdBy uint) (*model.SSHKey, error) { + name = strings.TrimSpace(name) + if name == "" || len(name) > 64 { + return nil, fmt.Errorf("%w: 密钥名称需为 1~64 字符", ErrKeyInvalid) + } + keyType, fingerprint, body, err := parsePublicKey(publicKey) + if err != nil { + return nil, err + } + u, err := s.user(ctx, userID) + if err != nil { + return nil, err + } + if u.Status != model.UserStatusActive { + return nil, ErrUserNotActive + } + // 同用户下已存在该公钥(active)→ 拒绝重复;已吊销的密钥允许重新上传 + var n int64 + if err := s.db.WithContext(ctx).Model(&model.SSHKey{}). + Where("user_id = ? AND fingerprint = ? AND status = ?", userID, fingerprint, model.StatusActive). + Count(&n).Error; err != nil { + return nil, err + } + if n > 0 { + return nil, ErrKeyDuplicate + } + k := &model.SSHKey{ + UserID: userID, + Name: name, + KeyType: keyType, + PublicKey: body, + Fingerprint: fingerprint, + Status: model.StatusActive, + Source: "user_uploaded", + CreatedBy: createdBy, + } + if err := s.db.WithContext(ctx).Create(k).Error; err != nil { + return nil, err + } + if err := s.syncUserKeys(ctx, u); err != nil { + _ = s.db.WithContext(ctx).Delete(k).Error + return nil, err + } + return k, nil +} + +// Rename 重命名密钥(仅元数据,不影响 authorized_keys)。 +func (s *KeyService) Rename(ctx context.Context, keyID, userID uint, name string) (*model.SSHKey, error) { + name = strings.TrimSpace(name) + if name == "" || len(name) > 64 { + return nil, fmt.Errorf("%w: 密钥名称需为 1~64 字符", ErrKeyInvalid) + } + k, err := s.owned(ctx, keyID, userID) + if err != nil { + return nil, err + } + if err := s.db.WithContext(ctx).Model(&model.SSHKey{}).Where("id = ?", keyID).Update("name", name).Error; err != nil { + return nil, err + } + k.Name = name + return k, nil +} + +// Revoke 吊销密钥(软删除,保留记录供审计)。先以剩余有效密钥(排除本次 +// 吊销的密钥)重写 authorized_keys(吊销立即失效),成功后再落 DB;同步失败 +// 则中止,文件与 DB 保持一致(密钥仍为 active)。已吊销时幂等。 +func (s *KeyService) Revoke(ctx context.Context, keyID, userID uint) (*model.SSHKey, error) { + k, err := s.owned(ctx, keyID, userID) + if err != nil { + return nil, err + } + if k.Status == model.StatusRevoked { + return k, nil // 幂等 + } + u, err := s.user(ctx, userID) + if err != nil { + return nil, err + } + // 剩余有效密钥(不含本次吊销的),先同步文件再落 DB(同 Disable 的 fail-closed 模式) + keys, err := activeUserKeys(s.db, ctx, u.ID, k.ID) + if err != nil { + return nil, err + } + if err := s.sys.SyncAuthorizedKeys(ctx, u.Username, keys); err != nil { + return nil, err + } + now := time.Now() + if err := s.db.WithContext(ctx).Model(&model.SSHKey{}).Where("id = ?", keyID). + Updates(map[string]any{"status": model.StatusRevoked, "revoked_at": &now}).Error; err != nil { + return nil, err + } + k.Status = model.StatusRevoked + k.RevokedAt = &now + return k, nil +} + +// owned 返回属于 userID 的密钥;跨用户访问视为不存在,不泄露存在性。 +func (s *KeyService) owned(ctx context.Context, keyID, userID uint) (*model.SSHKey, error) { + var k model.SSHKey + if err := s.db.WithContext(ctx).First(&k, "id = ? AND user_id = ?", keyID, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrKeyNotFound + } + return nil, err + } + return &k, nil +} + +// syncUserKeys 以 DB 当前 active 密钥全量重写 authorized_keys。 +func (s *KeyService) syncUserKeys(ctx context.Context, u *model.User) error { + keys, err := activeUserKeys(s.db, ctx, u.ID, 0) + if err != nil { + return err + } + return s.sys.SyncAuthorizedKeys(ctx, u.Username, keys) +} + +// parsePublicKey 校验并解析 OpenSSH 公钥行,返回类型 / SHA256 指纹 / base64 主体。 +// 仅接受标准单行 "类型 base64 [注释]";拒绝 ssh-dss(弱算法)与 <2048 位 RSA。 +func parsePublicKey(input string) (keyType, fingerprint, body string, err error) { + if len(input) > maxPublicKeyLen { + return "", "", "", fmt.Errorf("%w: 公钥内容过长", ErrKeyInvalid) + } + line := strings.TrimSpace(input) + if line == "" || strings.ContainsAny(line, "\r\n") { + return "", "", "", fmt.Errorf("%w: 公钥必须为单行", ErrKeyInvalid) + } + pub, _, options, rest, perr := ssh.ParseAuthorizedKey([]byte(line)) + if perr != nil { + return "", "", "", fmt.Errorf("%w: %v", ErrKeyInvalid, perr) + } + if len(options) > 0 || len(rest) > 0 { + return "", "", "", fmt.Errorf("%w: 仅支持标准公钥行,不能带选项或多余内容", ErrKeyInvalid) + } + keyType = pub.Type() + switch keyType { + case "ssh-dss": + return "", "", "", fmt.Errorf("%w: 不支持 ssh-dss 密钥", ErrKeyInvalid) + case "ssh-rsa": + if cp, ok := pub.(ssh.CryptoPublicKey); ok { + if rsaPub, ok := cp.CryptoPublicKey().(*rsa.PublicKey); ok && rsaPub.N.BitLen() < 2048 { + return "", "", "", fmt.Errorf("%w: RSA 密钥长度至少 2048 位", ErrKeyInvalid) + } + } + } + fingerprint = ssh.FingerprintSHA256(pub) + body = base64.StdEncoding.EncodeToString(pub.Marshal()) + return keyType, fingerprint, body, nil +} diff --git a/internal/service/key_test.go b/internal/service/key_test.go new file mode 100644 index 0000000..df073da --- /dev/null +++ b/internal/service/key_test.go @@ -0,0 +1,247 @@ +package service + +import ( + "context" + "crypto/dsa" + "crypto/ed25519" + "crypto/rand" + "crypto/rsa" + "errors" + "strings" + "testing" + "time" + + "golang.org/x/crypto/ssh" + + "ws_usernode/internal/model" +) + +// testPubKey 生成一条合法的 ed25519 公钥行。 +func testPubKey(t *testing.T) string { + t.Helper() + pub, _, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatalf("gen ed25519: %v", err) + } + sshPub, err := ssh.NewPublicKey(pub) + if err != nil { + t.Fatalf("ssh key: %v", err) + } + return strings.TrimSpace(string(ssh.MarshalAuthorizedKey(sshPub))) +} + +// testRSAKey 生成指定 bit 的 RSA 公钥行。 +func testRSAKey(t *testing.T, bits int) string { + t.Helper() + priv, err := rsa.GenerateKey(rand.Reader, bits) + if err != nil { + t.Fatalf("gen rsa: %v", err) + } + sshPub, err := ssh.NewPublicKey(&priv.PublicKey) + if err != nil { + t.Fatalf("ssh key: %v", err) + } + return strings.TrimSpace(string(ssh.MarshalAuthorizedKey(sshPub))) +} + +// testDSSKey 生成一条 ssh-dss 公钥行(弱算法,应被拒绝)。 +func testDSSKey(t *testing.T) string { + t.Helper() + var params dsa.Parameters + if err := dsa.GenerateParameters(¶ms, rand.Reader, dsa.L1024N160); err != nil { + t.Fatalf("dsa params: %v", err) + } + priv := new(dsa.PrivateKey) + priv.PublicKey.Parameters = params + if err := dsa.GenerateKey(priv, rand.Reader); err != nil { + t.Fatalf("dsa key: %v", err) + } + sshPub, err := ssh.NewPublicKey(&priv.PublicKey) + if err != nil { + t.Fatalf("ssh dsa key: %v", err) + } + return strings.TrimSpace(string(ssh.MarshalAuthorizedKey(sshPub))) +} + +func TestParsePublicKey(t *testing.T) { + ed := testPubKey(t) + keyType, fp, body, err := parsePublicKey(ed) + if err != nil { + t.Fatalf("parse valid ed25519: %v", err) + } + if keyType != "ssh-ed25519" { + t.Fatalf("keyType = %q, want ssh-ed25519", keyType) + } + if !strings.HasPrefix(fp, "SHA256:") || len(fp) != len("SHA256:")+43 { + t.Fatalf("fingerprint = %q", fp) + } + if body == "" { + t.Fatal("body empty") + } + // 同一输入解析结果稳定(指纹一致) + if _, fp2, _, err := parsePublicKey(ed); err != nil || fp2 != fp { + t.Fatalf("fingerprint not stable: %q vs %q err=%v", fp, fp2, err) + } + + // 合法 RSA-2048 + if _, _, _, err := parsePublicKey(testRSAKey(t, 2048)); err != nil { + t.Fatalf("parse rsa2048: %v", err) + } + // RSA-1024 拒绝 + if _, _, _, err := parsePublicKey(testRSAKey(t, 1024)); err == nil { + t.Fatal("rsa1024 should be rejected") + } + // ssh-dss 拒绝 + if _, _, _, err := parsePublicKey(testDSSKey(t)); err == nil { + t.Fatal("ssh-dss should be rejected") + } + // 多行 / 空 / 垃圾 / 选项 / 超长 + cases := []string{ + ed + "\n" + ed, + "", + "garbage not a key", + `command="echo x" ` + ed, + strings.Repeat("A", maxPublicKeyLen+1), + } + for _, in := range cases { + if _, _, _, err := parsePublicKey(in); err == nil { + t.Fatalf("input %q should be rejected", in[:min(len(in), 24)]) + } + } +} + +func TestKeyServiceLifecycle(t *testing.T) { + db := testDB(t) + sys := newFakeSys() + cfg := testConfig() + us := NewUserService(db, sys, cfg) + ks := NewKeyService(db, sys, cfg) + ctx := context.Background() + + u, err := us.Create(ctx, "wangwu", "ww@example.com", "", "", 0, 1) + if err != nil { + t.Fatalf("create user: %v", err) + } + pub := testPubKey(t) + + // 上传 → DB 落一条,authorized_keys 同步该密钥 + k, err := ks.Create(ctx, u.ID, "workstation", pub, u.ID) + if err != nil { + t.Fatalf("create key: %v", err) + } + if k.Fingerprint == "" || k.Status != model.StatusActive { + t.Fatalf("key fields: %+v", k) + } + synced := sys.keys["ext_wangwu"] + if len(synced) != 1 || synced[0].PublicKey == "" { + t.Fatalf("synced keys = %+v, want 1 key", synced) + } + + // 重复上传同一公钥 → 拒绝 + if _, err := ks.Create(ctx, u.ID, "dup", pub, u.ID); !errors.Is(err, ErrKeyDuplicate) { + t.Fatalf("duplicate err = %v, want ErrKeyDuplicate", err) + } + + // 列表 + keys, err := ks.ListByUser(ctx, u.ID) + if err != nil || len(keys) != 1 { + t.Fatalf("list keys = %v, err = %v", keys, err) + } + + // 重命名(不影响同步内容) + k2, err := ks.Rename(ctx, k.ID, u.ID, "home-laptop") + if err != nil { + t.Fatalf("rename: %v", err) + } + if k2.Name != "home-laptop" { + t.Fatalf("renamed = %q", k2.Name) + } + + // 吊销 → authorized_keys 清空(立即失效),DB 置 revoked + k3, err := ks.Revoke(ctx, k.ID, u.ID) + if err != nil { + t.Fatalf("revoke: %v", err) + } + if k3.Status != model.StatusRevoked || k3.RevokedAt == nil { + t.Fatalf("revoked key: %+v", k3) + } + if len(sys.keys["ext_wangwu"]) != 0 { + t.Fatalf("after revoke synced keys = %+v, want empty", sys.keys["ext_wangwu"]) + } + // 幂等 + if _, err := ks.Revoke(ctx, k.ID, u.ID); err != nil { + t.Fatalf("revoke again: %v", err) + } + + // 跨用户访问 → ErrKeyNotFound(不泄露存在性) + other, err := us.Create(ctx, "zhaoliu", "zl@example.com", "", "", 0, 1) + if err != nil { + t.Fatalf("create other user: %v", err) + } + if _, err := ks.Revoke(ctx, k.ID, other.ID); !errors.Is(err, ErrKeyNotFound) { + t.Fatalf("cross-user revoke err = %v, want ErrKeyNotFound", err) + } + if _, err := ks.Rename(ctx, k.ID, other.ID, "x"); !errors.Is(err, ErrKeyNotFound) { + t.Fatalf("cross-user rename err = %v, want ErrKeyNotFound", err) + } +} + +func TestKeyServiceCreateRollbackOnSyncFailure(t *testing.T) { + db := testDB(t) + sys := newFakeSys() + sys.syncErr = errors.New("sync boom") + cfg := testConfig() + us := NewUserService(db, sys, cfg) + ks := NewKeyService(db, sys, cfg) + ctx := context.Background() + + u, err := us.Create(ctx, "liuqian", "lq@example.com", "", "", 0, 1) + if err != nil { + t.Fatalf("create user: %v", err) + } + // 同步失败 → Create 报错且 DB 无残留记录 + if _, err := ks.Create(ctx, u.ID, "k", testPubKey(t), u.ID); err == nil { + t.Fatal("create with failing sync should error") + } + var n int64 + if err := db.Model(&model.SSHKey{}).Count(&n).Error; err != nil { + t.Fatalf("count: %v", err) + } + if n != 0 { + t.Fatalf("rollback failed: %d key rows remain", n) + } +} + +func TestKeyServiceCreateUserNotActive(t *testing.T) { + db := testDB(t) + sys := newFakeSys() + cfg := testConfig() + us := NewUserService(db, sys, cfg) + ks := NewKeyService(db, sys, cfg) + ctx := context.Background() + + u, err := us.Create(ctx, "sunqi", "sq@example.com", "", "", 0, 1) + if err != nil { + t.Fatalf("create user: %v", err) + } + if err := us.Disable(ctx, u.ID); err != nil { + t.Fatalf("disable: %v", err) + } + if _, err := ks.Create(ctx, u.ID, "k", testPubKey(t), u.ID); !errors.Is(err, ErrUserNotActive) { + t.Fatalf("create on disabled user err = %v, want ErrUserNotActive", err) + } + // 过期用户同样拒绝 + past := time.Now().Add(-time.Hour) + db.Model(&model.User{}).Where("id = ?", u.ID).Updates(map[string]any{"status": model.UserStatusExpired, "expire_at": &past}) + if _, err := ks.Create(ctx, u.ID, "k", testPubKey(t), u.ID); !errors.Is(err, ErrUserNotActive) { + t.Fatalf("create on expired user err = %v, want ErrUserNotActive", err) + } +} + +func TestKeyServiceListMissingUser(t *testing.T) { + db := testDB(t) + ks := NewKeyService(db, newFakeSys(), testConfig()) + if _, err := ks.ListByUser(context.Background(), 999); !errors.Is(err, ErrUserNotFound) { + t.Fatalf("list missing user err = %v, want ErrUserNotFound", err) + } +} diff --git a/internal/service/user.go b/internal/service/user.go index daeac1c..f48d2fa 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -211,8 +211,8 @@ func (s *UserService) Enable(ctx context.Context, id uint) error { if !s.systemAccountOK(ctx, u.Username) { return ErrSystemAccountMissing } - // 恢复有效密钥(M1 阶段用户尚无密钥,M2 接入后按 DB 同步) - keys, err := s.activeKeys(ctx, u.ID) + // 恢复有效密钥(以 DB 状态全量同步) + keys, err := activeUserKeys(s.db, ctx, u.ID, 0) if err != nil { return err } @@ -239,7 +239,7 @@ func (s *UserService) Extend(ctx context.Context, id uint, days int) error { if !s.systemAccountOK(ctx, u.Username) { return ErrSystemAccountMissing } - keys, err := s.activeKeys(ctx, u.ID) + keys, err := activeUserKeys(s.db, ctx, u.ID, 0) if err != nil { return err } @@ -271,14 +271,19 @@ func (s *UserService) Delete(ctx context.Context, id uint) error { }) } -// activeKeys 返回用户当前有效(active)密钥,供授权同步(M2 完善密钥管理)。 -func (s *UserService) activeKeys(ctx context.Context, userID uint) ([]system.Key, error) { +// activeUserKeys 返回用户当前有效(active)密钥,供 authorized_keys 全量同步 +// (UserService.Enable/Extend 与 KeyService 变更共用,保证同步口径一致)。 +// excludeKeyID 非 0 时排除指定密钥(吊销场景:先同步剩余密钥,再落 DB)。 +func activeUserKeys(db *gorm.DB, ctx context.Context, userID uint, excludeKeyID 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 { + if err := 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 { + if excludeKeyID != 0 && k.ID == excludeKeyID { + continue + } keys = append(keys, system.Key{Type: k.KeyType, PublicKey: k.PublicKey}) } return keys, nil diff --git a/internal/service/user_test.go b/internal/service/user_test.go index 119f3db..68f7321 100644 --- a/internal/service/user_test.go +++ b/internal/service/user_test.go @@ -20,6 +20,7 @@ type fakeSys struct { accounts map[string]bool keys map[string][]system.Key // username -> 最后一次同步的密钥 lastCmd string + syncErr error // 注入 SyncAuthorizedKeys 失败(测试回滚) } func newFakeSys() *fakeSys { @@ -58,6 +59,9 @@ func (f *fakeSys) Exists(_ context.Context, username string) (bool, error) { func (f *fakeSys) SyncAuthorizedKeys(_ context.Context, username string, keys []system.Key) error { f.mu.Lock() defer f.mu.Unlock() + if f.syncErr != nil { + return f.syncErr + } f.keys[username] = keys f.lastCmd = "sync-keys " + username return nil diff --git a/internal/system/system.go b/internal/system/system.go index 4f6c2f4..b5d1d57 100644 --- a/internal/system/system.go +++ b/internal/system/system.go @@ -17,6 +17,7 @@ import ( "context" "errors" "fmt" + "log/slog" "os" "os/exec" "path/filepath" @@ -57,26 +58,34 @@ type Manager interface { } // 本地实现注入的命令白名单(与 deploy/sudoers 保持一致)。 +// useradd/usermod/userdel/passwd 为账号生命周期;mkdir/chmod/chown/install +// 用于 authorized_keys 原子同步(M2):sudo 模式下经白名单命令落位并修正属主, +// 保证 sshd StrictModes 通过。 var allowedCommands = map[string]bool{ "useradd": true, "usermod": true, "userdel": true, "passwd": true, "chsh": true, + "mkdir": true, + "chmod": true, + "chown": true, + "install": true, } // localManager 为本地实现。 type localManager struct { cfg config.SystemConfig + log *slog.Logger } -// New 创建系统账号 Manager。 -func New(cfg config.SystemConfig) Manager { - return &localManager{cfg: cfg} +// New 创建系统账号 Manager。log 用于 dry-run 模式打印计划命令(演练提示)。 +func New(cfg config.SystemConfig, log *slog.Logger) Manager { + return &localManager{cfg: cfg, log: log} } // run 执行白名单命令。参数在调用处强校验。 -// 模式:dry-run 只返回命令文本;direct(sudo=false)直接执行; +// 模式:dry-run 只返回命令文本(并打印计划);direct(sudo=false)直接执行; // sudo 经 sudo -n 提权执行(对应 deploy/sudoers 白名单)。 func (m *localManager) run(ctx context.Context, cmd string, args ...string) (string, error) { if !allowedCommands[cmd] { @@ -84,6 +93,9 @@ func (m *localManager) run(ctx context.Context, cmd string, args ...string) (str } cmdline := strings.Join(append([]string{cmd}, args...), " ") if m.cfg.DryRun { + if m.log != nil { + m.log.Info("system: dry-run", "cmd", cmdline) + } return cmdline, nil } var ( @@ -165,16 +177,12 @@ func (m *localManager) SetLock(ctx context.Context, username string, locked bool // 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 +// passwdEntry 解析 /etc/passwd 中指定用户名的一行,返回 uid/gid。 +// 文件为世界可读,无需提权。 +func passwdEntry(username string) (uid, gid int, ok bool, err error) { f, err := os.Open(passwdPath) if err != nil { - return false, err + return 0, 0, false, err } defer f.Close() sc := bufio.NewScanner(f) @@ -183,32 +191,50 @@ func (m *localManager) Exists(ctx context.Context, username string) (bool, error if line == "" { continue } - fields := strings.SplitN(line, ":", 2) - if len(fields) > 0 && fields[0] == username { - return true, nil + fields := strings.Split(line, ":") + if len(fields) >= 4 && fields[0] == username { + uid, err1 := strconv.Atoi(fields[2]) + gid, err2 := strconv.Atoi(fields[3]) + if err1 != nil || err2 != nil { + return 0, 0, false, fmt.Errorf("system: parse passwd entry %q: %v/%v", username, err1, err2) + } + return uid, gid, true, nil } } if err := sc.Err(); err != nil { - return false, err + return 0, 0, false, err } - return false, nil + return 0, 0, false, nil } -// SyncAuthorizedKeys 全量重写 authorized_keys: +// 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 + _, _, ok, err := passwdEntry(username) + return ok, err +} + +// SyncAuthorizedKeys 全量重写 authorized_keys(PLAN F3:原子写 + 并发锁, +// 全量重写基于 DB 状态): // -// 1. 以 /.lock 文件锁串行化并发写(flock); -// 2. 写临时文件(0600),再 rename 原子替换; -// 3. 无有效密钥时写入空文件(SSH 行为一致,避免文件缺失导致的歧义)。 +// 1. 以 /usernode-keys-.lock 文件锁串行化并发写(flock); +// 2. 无有效密钥时写入空文件(SSH 行为一致,避免文件缺失导致的歧义); +// 3. sudo 模式经白名单命令 mkdir/chown/chmod/install 落位并修正属主 +// (sshd StrictModes 要求 .ssh 属用户且 0700);direct 模式进程直写, +// root 时同样修正属主。 // -// dry-run 模式只打印计划;真实写路径(direct/sudo)需要进程具备对家目录 -// .ssh 的写权限——生产部署经 sudoers 白名单授予受控脚本实现(M2 细化)。 +// dry-run 模式只打印计划;真实写路径(direct/sudo)见 deploy/sudoers.example。 func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, keys []Key) error { username = m.sysName(username) if err := pkg.ValidateSystemAccount(username); err != nil { return err } - home := filepath.Join(m.cfg.HomeBase, username) - sshDir := filepath.Join(home, m.cfg.AuthorizedKeysDir) + sshDir := filepath.Join(m.cfg.HomeBase, username, m.cfg.AuthorizedKeysDir) + dest := filepath.Join(sshDir, "authorized_keys") lockPath := filepath.Join(os.TempDir(), "usernode-keys-"+username+".lock") content := new(strings.Builder) @@ -220,18 +246,24 @@ func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, 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") + if m.cfg.Sudo { + plan.WriteString("sudo -n mkdir -p " + sshDir + " && sudo -n chown " + username + ":" + m.cfg.Group + " " + sshDir + " && sudo -n chmod 0700 " + sshDir + "\n") + plan.WriteString("write " + filepath.Join(os.TempDir(), "usernode-ak-"+username+".tmp") + " (0600)\n") + plan.WriteString("sudo -n install -o " + username + " -g " + m.cfg.Group + " -m 0600 " + dest + "\n") + } else { + plan.WriteString("mkdir -p " + sshDir + " (0700)\n") + plan.WriteString("flock " + lockPath + "\n") + plan.WriteString("write " + dest + " (0600, tmp+rename)\n") + } plan.WriteString(content.String()) // 演练模式:仅日志输出计划,不落盘 + if m.log != nil { + m.log.Info("system: dry-run plan", "account", username, "plan", "\n"+plan.String()) + } return nil } - // 生产路径:并发锁 + 原子写 - if err := os.MkdirAll(sshDir, 0o700); err != nil { - return err - } + // 并发锁串行化(进程内;跨节点由单实例部署保证,PLAN §11) lock, err := os.OpenFile(lockPath, os.O_CREATE|os.O_RDWR, 0o600) if err != nil { return err @@ -242,11 +274,47 @@ func (m *localManager) SyncAuthorizedKeys(ctx context.Context, username string, } defer funlock(lock) + if m.cfg.Sudo { + // .ssh 目录与属主:sshd StrictModes 要求属用户且 0700 + if _, err := m.run(ctx, "mkdir", "-p", sshDir); err != nil { + return err + } + if _, err := m.run(ctx, "chown", username+":"+m.cfg.Group, sshDir); err != nil { + return err + } + if _, err := m.run(ctx, "chmod", "0700", sshDir); err != nil { + return err + } + // 临时文件写入进程可写目录,经 install 原子落位(dest 目录内临时文件 + rename) + tmp := filepath.Join(os.TempDir(), "usernode-ak-"+username+"-"+strconv.Itoa(os.Getpid())) + if err := os.WriteFile(tmp, []byte(content.String()), 0o600); err != nil { + return err + } + defer os.Remove(tmp) + if _, err := m.run(ctx, "install", "-o", username, "-g", m.cfg.Group, "-m", "0600", tmp, dest); err != nil { + return err + } + return nil + } + + // direct 模式:进程直写(容器/测试用户验证);root 时修正属主 + if err := os.MkdirAll(sshDir, 0o700); err != nil { + return err + } + if os.Geteuid() == 0 { + if uid, gid, ok, err := passwdEntry(username); err != nil { + return err + } else if ok { + if err := os.Chown(sshDir, uid, gid); err != nil { + return err + } + } + } tmp := filepath.Join(sshDir, "authorized_keys.tmp."+strconv.Itoa(os.Getpid())) if err := os.WriteFile(tmp, []byte(content.String()), 0o600); err != nil { return err } - if err := os.Rename(tmp, filepath.Join(sshDir, "authorized_keys")); err != nil { + if err := os.Rename(tmp, dest); err != nil { os.Remove(tmp) return err } diff --git a/internal/system/system_test.go b/internal/system/system_test.go index 18a1982..da37e48 100644 --- a/internal/system/system_test.go +++ b/internal/system/system_test.go @@ -2,6 +2,8 @@ package system import ( "context" + "io" + "log/slog" "os" "path/filepath" "testing" @@ -9,6 +11,11 @@ import ( "ws_usernode/internal/config" ) +// discardLog 供测试注入的静默日志器。 +func discardLog() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + func testCfg() config.SystemConfig { cfg := config.Default() return cfg.System @@ -17,7 +24,7 @@ func testCfg() config.SystemConfig { func TestDryRunDoesNotExecute(t *testing.T) { cfg := testCfg() cfg.DryRun = true - m := New(cfg) + m := New(cfg, discardLog()) ctx := context.Background() // dry-run 下创建/删除不应报错(只打印计划) @@ -52,7 +59,7 @@ func TestExistsParsesPasswd(t *testing.T) { passwdPath = p t.Cleanup(func() { passwdPath = old }) - m := New(testCfg()) + m := New(testCfg(), discardLog()) ctx := context.Background() // 已带前缀与未带前缀都会命中同一账号 for _, name := range []string{"ext_zhangsan", "zhangsan"} { @@ -85,3 +92,68 @@ func TestPrefixNormalization(t *testing.T) { t.Fatalf("sysName(x) = %q", got) } } + +func TestSyncAuthorizedKeysDirect(t *testing.T) { + cfg := testCfg() + cfg.DryRun = false + cfg.Sudo = false + cfg.HomeBase = t.TempDir() // 临时家目录基路径,进程用户直写 + m := New(cfg, discardLog()) + ctx := context.Background() + + keys := []Key{ + {Type: "ssh-ed25519", PublicKey: "AAAAC3NzaC1lZDI1NTE5AAAAIB-test"}, + {Type: "ssh-rsa", PublicKey: "AAAAB3NzaC1yc2EAAAADAQAB-test"}, + } + if err := m.SyncAuthorizedKeys(ctx, "ext_zhangsan", keys); err != nil { + t.Fatalf("sync keys: %v", err) + } + sshDir := filepath.Join(cfg.HomeBase, "ext_zhangsan", ".ssh") + dest := filepath.Join(sshDir, "authorized_keys") + content, err := os.ReadFile(dest) + if err != nil { + t.Fatalf("read authorized_keys: %v", err) + } + want := "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIB-test\nssh-rsa AAAAB3NzaC1yc2EAAAADAQAB-test\n" + if string(content) != want { + t.Fatalf("authorized_keys = %q, want %q", content, want) + } + // .ssh 目录 0700、文件 0600 + di, err := os.Stat(sshDir) + if err != nil { + t.Fatalf("stat .ssh: %v", err) + } + if di.Mode().Perm() != 0o700 { + t.Fatalf(".ssh mode = %v, want 0700", di.Mode().Perm()) + } + fi, err := os.Stat(dest) + if err != nil { + t.Fatalf("stat authorized_keys: %v", err) + } + if fi.Mode().Perm() != 0o600 { + t.Fatalf("authorized_keys mode = %v, want 0600", fi.Mode().Perm()) + } + + // 无有效密钥 → 写空文件(SSH 行为一致) + if err := m.SyncAuthorizedKeys(ctx, "ext_zhangsan", nil); err != nil { + t.Fatalf("sync empty: %v", err) + } + content, err = os.ReadFile(dest) + if err != nil { + t.Fatalf("re-read authorized_keys: %v", err) + } + if len(content) != 0 { + t.Fatalf("empty sync left content: %q", content) + } +} + +func TestSyncAuthorizedKeysDryRunSudoPlan(t *testing.T) { + cfg := testCfg() + cfg.DryRun = true + cfg.Sudo = true + m := New(cfg, discardLog()) + // sudo 模式的 dry-run 只打印计划(install 落位路径),不执行任何命令 + if err := m.SyncAuthorizedKeys(context.Background(), "ext_zhangsan", []Key{{Type: "ssh-ed25519", PublicKey: "AAA"}}); err != nil { + t.Fatalf("dry-run sudo plan: %v", err) + } +}