package app import ( "bytes" "context" "database/sql" "database/sql/driver" "encoding/json" "fmt" "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/go-sql-driver/mysql" ) // This transactional fixture exercises the real HTTP handlers without needing // production credentials. Unexpected SQL is an error, including session writes // during administrator provisioning. type provisionedUser struct { publicID, passwordHash, nickname, city, bio string phoneHash, phoneCipher []byte gender int64 privacy, notification bool audit []byte sessions int } type provisioningDB struct { t *testing.T user, pending *provisionedUser allowed, inTx bool failAt string failure error begins, rollbacks int } type provisioningConnector struct{ store *provisioningDB } func (c provisioningConnector) Connect(context.Context) (driver.Conn, error) { return c.store, nil } func (provisioningConnector) Driver() driver.Driver { return oauthTestDriver{} } func (*provisioningDB) Prepare(string) (driver.Stmt, error) { return nil, fmt.Errorf("unexpected prepare") } func (*provisioningDB) Close() error { return nil } func (s *provisioningDB) Begin() (driver.Tx, error) { s.begins++ s.inTx = true s.pending = nil if s.user != nil { copy := *s.user s.pending = © } return s, nil } func (s *provisioningDB) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { return s.Begin() } func (s *provisioningDB) Commit() error { s.user, s.pending, s.inTx = s.pending, nil, false return nil } func (s *provisioningDB) Rollback() error { s.pending, s.inTx = nil, false s.rollbacks++ return nil } type provisioningInsertResult struct{} func (provisioningInsertResult) LastInsertId() (int64, error) { return 42, nil } func (provisioningInsertResult) RowsAffected() (int64, error) { return 1, nil } func (s *provisioningDB) ExecContext(_ context.Context, query string, args []driver.NamedValue) (driver.Result, error) { if s.failAt != "" && strings.Contains(query, s.failAt) { return nil, s.failure } if strings.Contains(query, "api_rate_limits") { return driver.RowsAffected(1), nil } if !s.inTx { s.t.Errorf("account mutation outside transaction: %s", query) return nil, fmt.Errorf("no transaction") } switch { case strings.HasPrefix(query, "INSERT INTO users "): if s.pending != nil && bytes.Equal(s.pending.phoneHash, args[1].Value.([]byte)) { return nil, &mysql.MySQLError{Number: 1062, Message: "Duplicate entry for key 'users.uk_users_phone_hash'"} } s.pending = &provisionedUser{publicID: args[0].Value.(string), phoneHash: args[1].Value.([]byte), phoneCipher: args[2].Value.([]byte), passwordHash: args[3].Value.(string)} return provisioningInsertResult{}, nil case strings.HasPrefix(query, "INSERT INTO user_profiles "): s.pending.nickname, s.pending.gender = args[1].Value.(string), args[2].Value.(int64) s.pending.city, s.pending.bio = args[3].Value.(string), args[4].Value.(string) case strings.HasPrefix(query, "INSERT INTO user_privacy_settings "): s.pending.privacy = true case strings.HasPrefix(query, "INSERT INTO user_notification_settings "): s.pending.notification = true case strings.HasPrefix(query, "INSERT INTO admin_audit_logs"): if args[0].Value != int64(7) || args[1].Value != int64(42) { s.t.Error("audit must identify both the administrator and created user") } s.pending.audit = args[2].Value.([]byte) case strings.HasPrefix(query, "INSERT INTO user_sessions "): s.pending.sessions++ case strings.HasPrefix(query, "UPDATE user_sessions "), strings.HasPrefix(query, "INSERT INTO user_devices"), strings.HasPrefix(query, "UPDATE user_profiles SET last_active_at"): default: s.t.Errorf("unexpected exec: %s", query) return nil, fmt.Errorf("unexpected exec") } return driver.RowsAffected(1), nil } func (s *provisioningDB) QueryContext(_ context.Context, query string, args []driver.NamedValue) (driver.Rows, error) { switch { case strings.Contains(query, "SELECT status,token_version FROM admin_users"): return oauthRow(int64(1), int64(0)), nil case strings.Contains(query, "SELECT token_version FROM"): return oauthRow(int64(0)), nil case strings.Contains(query, "FROM admin_user_roles"): allowed := int64(0) if s.allowed && args[1].Value == "users:create" { allowed = 1 } return oauthRow(allowed), nil case strings.Contains(query, "SELECT hits FROM api_rate_limits"): return oauthRow(int64(1)), nil case strings.HasPrefix(query, "SELECT u.id,u.password_hash,u.status,p.nickname"): if s.user == nil || !bytes.Equal(s.user.phoneHash, args[0].Value.([]byte)) { return &oauthTestRows{columns: []string{"id", "password_hash", "status", "nickname"}}, nil } return oauthRow(int64(42), s.user.passwordHash, int64(1), s.user.nickname), nil case strings.HasPrefix(query, "SELECT COUNT(*) FROM users u JOIN user_profiles"): return oauthRow(int64(1)), nil case strings.HasPrefix(query, "SELECT u.id,u.public_id,u.phone_cipher"): u := s.user return oauthRow(int64(42), u.publicID, u.phoneCipher, int64(1), int64(0), time.Now(), false, "", u.nickname, "", u.gender, u.city, int64(0), int64(0), nil, "UNVERIFIED", nil), nil default: s.t.Errorf("unexpected query: %s", query) return nil, fmt.Errorf("unexpected query") } } func provisioningApp(t *testing.T) (*App, *provisioningDB, http.HandlerFunc, string) { t.Helper() store := &provisioningDB{t: t, allowed: true} db := sql.OpenDB(provisioningConnector{store}) t.Cleanup(func() { _ = db.Close() }) a := &App{db: db, config: Config{JWTSecret: "test-provisioning-secret", Environment: "development"}} token, err := a.token(7, "admin", "tester", time.Hour) if err != nil { t.Fatal(err) } for _, route := range a.adminRoutes() { if route.Method == http.MethodPost && route.Path == "/admin/v1/users" { return a, store, route.Handler, token } } t.Fatal("admin create user route is not registered") return nil, nil, nil, "" } func provisionRequest(handler http.HandlerFunc, token string, payload any) *httptest.ResponseRecorder { body, _ := json.Marshal(payload) r := httptest.NewRequest(http.MethodPost, "/admin/v1/users", bytes.NewReader(body)) if token != "" { r.Header.Set("Authorization", "Bearer "+token) } w := httptest.NewRecorder() handler(w, r) return w } func validProvisionPayload() map[string]any { return map[string]any{"phone": "13800138000", "password": "Password123!", "nickname": " 管理员创建用户 ", "gender": 1, "city": " 北京 ", "bio": " 简介 "} } func TestAdminCreateUserAndPasswordLogin(t *testing.T) { a, store, handler, token := provisioningApp(t) w := provisionRequest(handler, token, validProvisionPayload()) if w.Code != http.StatusOK || store.user == nil { t.Fatalf("create failed: %d %s", w.Code, w.Body.String()) } u := store.user if !u.privacy || !u.notification || len(u.audit) == 0 || u.sessions != 0 { t.Fatal("account must have default settings and an audit, but no login session") } if u.nickname != "管理员创建用户" || u.city != "北京" || u.bio != "简介" { t.Fatal("profile whitespace was not normalized") } if len(u.publicID) > 20 || !strings.HasPrefix(u.publicID, "XY") { t.Fatal("invalid public ID") } phone, err := a.decryptPhone(u.phoneCipher) if err != nil || phone != "13800138000" || bytes.Contains(u.phoneCipher, []byte(phone)) { t.Fatal("phone was not encrypted correctly") } if !checkPassword(u.passwordHash, "Password123!") || u.passwordHash == "Password123!" { t.Fatal("password must be bcrypt hashed") } for _, sensitive := range []string{"13800138000", "Password123!", "accessToken", "refreshToken"} { if strings.Contains(w.Body.String(), sensitive) || bytes.Contains(u.audit, []byte(sensitive)) { t.Fatalf("creation leaked %s", sensitive) } } list := httptest.NewRecorder() a.adminUsers(list, httptest.NewRequest(http.MethodGet, "/admin/v1/users?keyword="+u.publicID, nil)) if list.Code != 200 || !strings.Contains(list.Body.String(), u.publicID) || !strings.Contains(list.Body.String(), `"phone":"138****8000"`) { t.Fatalf("new account missing from list: %s", list.Body.String()) } wrong := provisionRequest(a.loginPassword, "", map[string]any{"phone": phone, "password": "wrong-password"}) if wrong.Code != http.StatusUnauthorized || store.user.sessions != 0 { t.Fatal("incorrect initial password was accepted") } login := provisionRequest(a.loginPassword, "", map[string]any{"phone": phone, "password": "Password123!", "deviceId": "app-test"}) var result struct { Data struct { AccessToken string `json:"accessToken"` } `json:"data"` } if login.Code != http.StatusOK || json.Unmarshal(login.Body.Bytes(), &result) != nil || store.user.sessions != 1 { t.Fatalf("password login failed: %d %s", login.Code, login.Body.String()) } who, err := a.parseToken(result.Data.AccessToken) if err != nil || who.ID != 42 || who.Role != "user" { t.Fatal("login did not issue a valid user token") } } func TestAdminCreateUserPermissions(t *testing.T) { a, store, handler, token := provisioningApp(t) store.allowed = false if w := provisionRequest(handler, token, validProvisionPayload()); w.Code != http.StatusForbidden { t.Fatalf("missing create permission accepted: %d", w.Code) } if w := provisionRequest(handler, "", validProvisionPayload()); w.Code != http.StatusUnauthorized { t.Fatal("anonymous creation accepted") } userToken, _ := a.token(42, "user", "user", time.Hour) if w := provisionRequest(handler, userToken, validProvisionPayload()); w.Code != http.StatusUnauthorized { t.Fatal("client token accepted") } if store.begins != 0 { t.Fatal("unauthorized caller reached account creation") } } func TestAdminCreateUserValidation(t *testing.T) { cases := []struct { key string value any }{ {"phone", "12345678901"}, {"phone", ""}, {"password", ""}, {"nickname", " "}, {"nickname", strings.Repeat("名", 51)}, {"gender", 3}, {"gender", -1}, {"gender", 1.5}, {"city", strings.Repeat("城", 51)}, {"bio", strings.Repeat("文", 501)}, {"isTest", true}, {"status", 3}, {"vip", true}, } for i, tc := range cases { t.Run(fmt.Sprintf("%s-%d", tc.key, i), func(t *testing.T) { _, store, handler, token := provisioningApp(t) payload := validProvisionPayload() payload[tc.key] = tc.value w := provisionRequest(handler, token, payload) if w.Code != http.StatusBadRequest || store.begins != 0 { t.Fatalf("invalid input accepted: %d %s", w.Code, w.Body.String()) } }) } } func TestAdminCreateUserDuplicatePhonePreservesAccount(t *testing.T) { _, store, handler, token := provisioningApp(t) if w := provisionRequest(handler, token, validProvisionPayload()); w.Code != http.StatusOK { t.Fatal(w.Body.String()) } original := store.user payload := validProvisionPayload() payload["nickname"], payload["password"], payload["phone"] = "覆盖用户", "Different123", " 13800138000 " w := provisionRequest(handler, token, payload) if w.Code != http.StatusConflict || store.user != original || store.rollbacks != 1 { t.Fatalf("duplicate did not preserve original account: %d %s", w.Code, w.Body.String()) } } // Admin-set passwords follow the same rule as self-service ones, and whatever // passes it has to survive the round trip to a working login. func TestAdminCreatedUserCanLoginWithVariedPasswords(t *testing.T) { for _, password := range []string{"letters8", "中文密码八个字符", "Str0ng!Passw0rd", strings.Repeat("长密码", 30)} { a, store, handler, token := provisioningApp(t) payload := validProvisionPayload() payload["password"] = password created := provisionRequest(handler, token, payload) if created.Code != http.StatusOK { t.Fatalf("password rejected: %d %s", created.Code, created.Body.String()) } login := provisionRequest(a.loginPassword, "", map[string]any{"phone": payload["phone"], "password": password}) if login.Code != http.StatusOK || store.user.sessions != 1 { t.Fatalf("new password cannot login: %d %s", login.Code, login.Body.String()) } } } func TestAdminCreateUserRollsBackEveryFailedWrite(t *testing.T) { for _, table := range []string{"users", "user_profiles", "user_privacy_settings", "user_notification_settings", "admin_audit_logs"} { t.Run(table, func(t *testing.T) { _, store, handler, token := provisioningApp(t) store.failAt, store.failure = "INSERT INTO "+table, fmt.Errorf("database unavailable") w := provisionRequest(handler, token, validProvisionPayload()) if w.Code != http.StatusInternalServerError || store.user != nil || store.rollbacks != 1 { t.Fatalf("failed %s write left partial account: %d %s", table, w.Code, w.Body.String()) } if strings.Contains(w.Body.String(), "手机号已被使用") { t.Fatal("database error misreported as duplicate phone") } }) } }