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 syncErr error // 注入 SyncAuthorizedKeys 失败(测试回滚) } 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() if f.syncErr != nil { return f.syncErr } 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) } }