Files
kefu/im/backend/internal/app/admin_create_user_test.go
T
2026-09-03 08:38:17 +08:00

318 lines
13 KiB
Go

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 = &copy
}
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())
}
}
func TestAdminCreatedUserCanLoginWithSimpleOrLongPassword(t *testing.T) {
for _, password := range []string{"1", "letters", "中文", 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")
}
})
}
}