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) } }