feat(M2): SSH 密钥管理 — 公钥上传/重命名/吊销、authorized_keys 原子同步与吊销即时失效
- KeyService:crypto/ssh 解析校验(单行/类型/长度/去重指纹,拒 ssh-dss 与 RSA<2048), Create/Rename/Revoke/List,变更后以 DB 状态全量重写 authorized_keys(同步失败回滚) - system 层:SyncAuthorizedKeys 完善 —— sudo 模式经白名单命令(mkdir/chown/chmod/install) 落位并修正属主(sshd StrictModes),direct 模式 root 时同样修正属主;dry-run 计划日志 - API:GET/POST /me/keys、PATCH/DELETE /me/keys/:id(user 会话)、GET /users/:id/keys(admin), 密钥操作带审计;deploy/sudoers.example 补充密钥同步白名单 - 版本 0.3.0-m2;测试:service 单元(校验/生命周期/回滚/权限)、system 直写落盘、 API 全流程集成;容器 E2E 32 项 PASS(真实 useradd/authorized_keys/吊销即时失效/禁用清空/删除回收)
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user