530 lines
20 KiB
Go
530 lines
20 KiB
Go
package app
|
|
|
|
import (
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"database/sql"
|
|
"encoding/binary"
|
|
"fmt"
|
|
"net/http"
|
|
"net/url"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
type authRequest struct {
|
|
Phone string `json:"phone"`
|
|
Password string `json:"password"`
|
|
Code string `json:"code"`
|
|
Nickname string `json:"nickname"`
|
|
DeviceID string `json:"deviceId"`
|
|
Scene string `json:"scene"`
|
|
}
|
|
|
|
func (a *App) sendSMS(w http.ResponseWriter, r *http.Request) {
|
|
var req authRequest
|
|
if err := decode(r, &req); err != nil || !validPhone(req.Phone) {
|
|
fail(w, http.StatusBadRequest, 20001, "请输入正确的手机号")
|
|
return
|
|
}
|
|
if !a.configBool(r.Context(), "sms.enabled", true) {
|
|
fail(w, http.StatusServiceUnavailable, 50002, "短信服务暂未开放")
|
|
return
|
|
}
|
|
if req.Scene == "" {
|
|
req.Scene = "login"
|
|
}
|
|
if req.Scene != "login" && req.Scene != "register" && req.Scene != "reset" && req.Scene != "change_phone" {
|
|
fail(w, http.StatusBadRequest, 20001, "验证码场景无效")
|
|
return
|
|
}
|
|
phone := strings.TrimSpace(req.Phone)
|
|
if !a.rateLimit(w, r, "sms_ip", clientIP(r), 20, time.Hour) || !a.rateLimit(w, r, "sms_phone", phone, 5, time.Hour) {
|
|
return
|
|
}
|
|
var recent int
|
|
_ = a.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM sms_verification_codes WHERE phone_hash=? AND scene=? AND created_at>DATE_SUB(NOW(3),INTERVAL 60 SECOND)`, phoneHash(req.Phone), req.Scene).Scan(&recent)
|
|
if recent > 0 {
|
|
fail(w, http.StatusTooManyRequests, 20002, "请稍后再获取验证码")
|
|
return
|
|
}
|
|
provider := a.configPlain(r.Context(), "sms.provider", "debug")
|
|
if a.config.Environment == "production" && provider == "debug" {
|
|
fail(w, http.StatusServiceUnavailable, 50002, "生产环境禁止使用调试短信服务")
|
|
return
|
|
}
|
|
code := a.configPlain(r.Context(), "sms.debug_code", "123456")
|
|
if provider != "debug" {
|
|
var buffer [4]byte
|
|
_, _ = rand.Read(buffer[:])
|
|
code = fmt.Sprintf("%06d", binary.BigEndian.Uint32(buffer[:])%1_000_000)
|
|
}
|
|
if err := a.dispatchSMS(r.Context(), phone, req.Scene, code); err != nil {
|
|
fail(w, http.StatusBadGateway, 50002, err.Error())
|
|
return
|
|
}
|
|
expires, _ := strconv.Atoi(a.configPlain(r.Context(), "sms.expire_seconds", "300"))
|
|
if expires < 60 || expires > 1800 {
|
|
expires = 300
|
|
}
|
|
codeHash := sha256.Sum256([]byte(code))
|
|
_, err := a.db.ExecContext(r.Context(), `INSERT INTO sms_verification_codes(phone_hash,scene,code_hash,expires_at)VALUES(?,?,?,DATE_ADD(NOW(3),INTERVAL ? SECOND))`, phoneHash(req.Phone), req.Scene, codeHash[:], expires)
|
|
if err != nil {
|
|
fail(w, 500, 50001, "保存验证码失败")
|
|
return
|
|
}
|
|
data := map[string]any{"expiresIn": expires, "provider": provider}
|
|
if provider == "debug" {
|
|
data["debugCode"] = code
|
|
}
|
|
reply(w, data)
|
|
}
|
|
|
|
func (a *App) register(w http.ResponseWriter, r *http.Request) {
|
|
var req authRequest
|
|
if err := decode(r, &req); err != nil {
|
|
fail(w, http.StatusBadRequest, 20001, err.Error())
|
|
return
|
|
}
|
|
if !validPhone(req.Phone) || !validUserPassword(req.Password) || strings.TrimSpace(req.Nickname) == "" || len([]rune(strings.TrimSpace(req.Nickname))) > 50 {
|
|
fail(w, http.StatusBadRequest, 20001, "请填写有效的手机号、昵称和密码")
|
|
return
|
|
}
|
|
if !a.rateLimit(w, r, "register_ip", clientIP(r), 20, 10*time.Minute) || !a.rateLimit(w, r, "register_phone", strings.TrimSpace(req.Phone), 10, 10*time.Minute) {
|
|
return
|
|
}
|
|
if !a.consumeSMSCode(r, req.Phone, "register", req.Code) {
|
|
fail(w, http.StatusBadRequest, 20001, "验证码错误或已过期")
|
|
return
|
|
}
|
|
hash, err := hashPassword(req.Password)
|
|
if err != nil {
|
|
fail(w, 500, 50001, "创建账号失败")
|
|
return
|
|
}
|
|
tx, err := a.db.BeginTx(r.Context(), nil)
|
|
if err != nil {
|
|
fail(w, 500, 50001, "创建账号失败")
|
|
return
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
publicID := fmt.Sprintf("XY%d%s", time.Now().UnixMilli(), randomToken()[:5])
|
|
phoneCipher, encryptErr := a.encryptPhone(req.Phone)
|
|
if encryptErr != nil {
|
|
fail(w, 500, 50001, "加密账号信息失败")
|
|
return
|
|
}
|
|
result, err := tx.ExecContext(r.Context(), `INSERT INTO users (public_id,country_code,phone_hash,phone_cipher,password_hash) VALUES (?,'+86',?,?,?)`, publicID, phoneHash(req.Phone), phoneCipher, hash)
|
|
if err != nil {
|
|
fail(w, http.StatusConflict, 20001, "该手机号已注册")
|
|
return
|
|
}
|
|
userID, err := result.LastInsertId()
|
|
if err != nil {
|
|
fail(w, 500, 50001, "创建账号失败")
|
|
return
|
|
}
|
|
_, err = tx.ExecContext(r.Context(), `INSERT INTO user_profiles (user_id,nickname,bio,profile_score,last_active_at) VALUES (?,?, '遇见更好的陌生人',30,NOW(3))`, userID, req.Nickname)
|
|
if err != nil {
|
|
fail(w, 500, 50001, "创建资料失败")
|
|
return
|
|
}
|
|
if _, err = tx.ExecContext(r.Context(), `INSERT INTO user_privacy_settings (user_id) VALUES (?)`, userID); err != nil {
|
|
fail(w, 500, 50001, "创建账号失败")
|
|
return
|
|
}
|
|
if _, err = tx.ExecContext(r.Context(), `INSERT INTO user_notification_settings (user_id) VALUES (?)`, userID); err != nil {
|
|
fail(w, 500, 50001, "创建账号失败")
|
|
return
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
fail(w, 500, 50001, "创建账号失败")
|
|
return
|
|
}
|
|
a.finishLogin(w, r, userID, req.Nickname, req.DeviceID)
|
|
}
|
|
|
|
func (a *App) loginPassword(w http.ResponseWriter, r *http.Request) {
|
|
var req authRequest
|
|
if err := decode(r, &req); err != nil {
|
|
fail(w, 400, 20001, err.Error())
|
|
return
|
|
}
|
|
if !validPhone(req.Phone) || req.Password == "" {
|
|
fail(w, http.StatusUnauthorized, 10001, "手机号或密码错误")
|
|
return
|
|
}
|
|
if !a.rateLimit(w, r, "login_ip", clientIP(r), 60, 10*time.Minute) || !a.rateLimit(w, r, "login_phone", strings.TrimSpace(req.Phone), 10, 10*time.Minute) {
|
|
return
|
|
}
|
|
var id int64
|
|
var hash, nickname string
|
|
var status int
|
|
err := a.db.QueryRowContext(r.Context(), `SELECT u.id,u.password_hash,u.status,p.nickname FROM users u JOIN user_profiles p ON p.user_id=u.id WHERE u.phone_hash=? AND u.deleted_at IS NULL`, phoneHash(req.Phone)).Scan(&id, &hash, &status, &nickname)
|
|
if err != nil || !checkPassword(hash, req.Password) {
|
|
fail(w, http.StatusUnauthorized, 10001, "手机号或密码错误")
|
|
return
|
|
}
|
|
if status != 1 {
|
|
fail(w, http.StatusForbidden, 10006, "账号当前不可用")
|
|
return
|
|
}
|
|
a.finishLogin(w, r, id, nickname, req.DeviceID)
|
|
}
|
|
|
|
func (a *App) loginSMS(w http.ResponseWriter, r *http.Request) {
|
|
var req authRequest
|
|
if err := decode(r, &req); err != nil || !validPhone(req.Phone) || len(req.Code) != 6 {
|
|
fail(w, 400, 20001, "验证码格式错误")
|
|
return
|
|
}
|
|
if !a.rateLimit(w, r, "sms_login_ip", clientIP(r), 30, 10*time.Minute) || !a.rateLimit(w, r, "sms_login_phone", strings.TrimSpace(req.Phone), 10, 10*time.Minute) {
|
|
return
|
|
}
|
|
if !a.consumeSMSCode(r, req.Phone, "login", req.Code) {
|
|
fail(w, 400, 20001, "验证码错误或已过期")
|
|
return
|
|
}
|
|
var id int64
|
|
var nickname string
|
|
if err := a.db.QueryRowContext(r.Context(), `SELECT u.id,p.nickname FROM users u JOIN user_profiles p ON p.user_id=u.id WHERE u.phone_hash=? AND u.status=1`, phoneHash(req.Phone)).Scan(&id, &nickname); err != nil {
|
|
fail(w, http.StatusUnauthorized, 10001, "账号不存在")
|
|
return
|
|
}
|
|
a.finishLogin(w, r, id, nickname, req.DeviceID)
|
|
}
|
|
|
|
func (a *App) resetPassword(w http.ResponseWriter, r *http.Request) {
|
|
var req authRequest
|
|
if err := decode(r, &req); err != nil || !validPhone(req.Phone) || len(req.Code) != 6 || !validUserPassword(req.Password) {
|
|
fail(w, 400, 20001, "请填写有效的手机号、验证码和新密码")
|
|
return
|
|
}
|
|
if !a.rateLimit(w, r, "password_reset_ip", clientIP(r), 20, 10*time.Minute) || !a.rateLimit(w, r, "password_reset_phone", strings.TrimSpace(req.Phone), 10, 10*time.Minute) {
|
|
return
|
|
}
|
|
if !a.consumeSMSCode(r, req.Phone, "reset", req.Code) {
|
|
fail(w, 400, 20001, "验证码错误或已过期")
|
|
return
|
|
}
|
|
hash, err := hashPassword(req.Password)
|
|
if err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "密码加密失败")
|
|
return
|
|
}
|
|
result, err := a.db.ExecContext(r.Context(), `UPDATE users SET password_hash=? WHERE phone_hash=? AND deleted_at IS NULL`, hash, phoneHash(req.Phone))
|
|
if err != nil {
|
|
fail(w, 500, 50001, "重置密码失败")
|
|
return
|
|
}
|
|
affected, _ := result.RowsAffected()
|
|
if affected == 0 {
|
|
fail(w, 404, 30001, "账号不存在")
|
|
return
|
|
}
|
|
_, _ = a.db.ExecContext(r.Context(), `UPDATE user_sessions SET revoked_at=NOW(3) WHERE user_id IN (SELECT id FROM users WHERE phone_hash=?) AND revoked_at IS NULL`, phoneHash(req.Phone))
|
|
_, _ = a.db.ExecContext(r.Context(), `INSERT INTO user_security_controls(user_id,token_version,force_logout_at,password_reset_at) SELECT id,1,NOW(3),NOW(3) FROM users WHERE phone_hash=? ON DUPLICATE KEY UPDATE token_version=token_version+1,force_logout_at=VALUES(force_logout_at),password_reset_at=VALUES(password_reset_at)`, phoneHash(req.Phone))
|
|
var resetUserID int64
|
|
if a.db.QueryRowContext(r.Context(), `SELECT id FROM users WHERE phone_hash=?`, phoneHash(req.Phone)).Scan(&resetUserID) == nil {
|
|
a.hub.disconnect(resetUserID)
|
|
}
|
|
reply(w, map[string]bool{"success": true})
|
|
}
|
|
|
|
func (a *App) consumeSMSCode(r *http.Request, phone, scene, code string) bool {
|
|
if phone == "" || code == "" {
|
|
return false
|
|
}
|
|
var id int64
|
|
var expected []byte
|
|
err := a.db.QueryRowContext(r.Context(), `SELECT id,code_hash FROM sms_verification_codes WHERE phone_hash=? AND scene=? AND used_at IS NULL AND expires_at>NOW(3) ORDER BY id DESC LIMIT 1`, phoneHash(phone), scene).Scan(&id, &expected)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
actual := sha256.Sum256([]byte(code))
|
|
if !bytesEqual(expected, actual[:]) {
|
|
return false
|
|
}
|
|
result, err := a.db.ExecContext(r.Context(), `UPDATE sms_verification_codes SET used_at=NOW(3) WHERE id=? AND used_at IS NULL`, id)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
affected, _ := result.RowsAffected()
|
|
return affected == 1
|
|
}
|
|
|
|
func bytesEqual(left, right []byte) bool {
|
|
if len(left) != len(right) {
|
|
return false
|
|
}
|
|
var different byte
|
|
for index := range left {
|
|
different |= left[index] ^ right[index]
|
|
}
|
|
return different == 0
|
|
}
|
|
|
|
func (a *App) finishLogin(w http.ResponseWriter, r *http.Request, id int64, nickname, deviceID string) {
|
|
if deviceID == "" {
|
|
deviceID = "web-h5"
|
|
}
|
|
if len(deviceID) > 100 {
|
|
fail(w, http.StatusBadRequest, 20001, "设备标识过长")
|
|
return
|
|
}
|
|
accessToken, err := a.token(id, "user", nickname, 30*time.Minute)
|
|
if err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "创建登录令牌失败")
|
|
return
|
|
}
|
|
refresh := randomToken()
|
|
refreshHash := sha256.Sum256([]byte(refresh))
|
|
tx, err := a.db.BeginTx(r.Context(), &sql.TxOptions{Isolation: sql.LevelReadCommitted})
|
|
if err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "创建登录会话失败")
|
|
return
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
if _, err = tx.ExecContext(r.Context(), `UPDATE user_sessions SET revoked_at=NOW(3) WHERE user_id=? AND device_id=? AND revoked_at IS NULL`, id, deviceID); err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "更新设备会话失败")
|
|
return
|
|
}
|
|
if _, err = tx.ExecContext(r.Context(), `INSERT INTO user_sessions (user_id,device_id,refresh_token_hash,expires_at) VALUES (?,?,?,DATE_ADD(NOW(3), INTERVAL 60 DAY))`, id, deviceID, refreshHash[:]); err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "创建登录会话失败")
|
|
return
|
|
}
|
|
platform := strings.TrimSpace(r.Header.Get("X-Device-Platform"))
|
|
model := strings.TrimSpace(r.Header.Get("X-Device-Model"))
|
|
osVersion := strings.TrimSpace(r.Header.Get("X-OS-Version"))
|
|
if decoded, decodeErr := url.QueryUnescape(model); decodeErr == nil {
|
|
model = decoded
|
|
}
|
|
if decoded, decodeErr := url.QueryUnescape(osVersion); decodeErr == nil {
|
|
osVersion = decoded
|
|
}
|
|
appVersion := strings.TrimSpace(r.Header.Get("X-App-Version"))
|
|
if platform == "" {
|
|
platform = "unknown"
|
|
}
|
|
if len(platform) > 20 || len(model) > 100 || len(osVersion) > 50 || len(appVersion) > 30 {
|
|
fail(w, http.StatusBadRequest, 20001, "设备信息格式错误")
|
|
return
|
|
}
|
|
if _, err = tx.ExecContext(r.Context(), `INSERT INTO user_devices(user_id,device_id,platform,device_model,os_version,app_version,last_ip,last_active_at,status)
|
|
VALUES(?,?,?,?,?,?,?,NOW(3),1)
|
|
ON DUPLICATE KEY UPDATE platform=VALUES(platform),device_model=VALUES(device_model),os_version=VALUES(os_version),app_version=VALUES(app_version),last_ip=VALUES(last_ip),last_active_at=NOW(3),status=1`, id, deviceID, platform, model, osVersion, appVersion, clientIP(r)); err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "保存设备信息失败")
|
|
return
|
|
}
|
|
if _, err = tx.ExecContext(r.Context(), `UPDATE user_profiles SET last_active_at=NOW(3) WHERE user_id=?`, id); err != nil || tx.Commit() != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "保存登录状态失败")
|
|
return
|
|
}
|
|
reply(w, map[string]any{"accessToken": accessToken, "refreshToken": refresh, "expiresIn": 1800, "userId": id})
|
|
}
|
|
|
|
func (a *App) refreshToken(w http.ResponseWriter, r *http.Request) {
|
|
var req struct {
|
|
RefreshToken string `json:"refreshToken"`
|
|
}
|
|
if err := decode(r, &req); err != nil {
|
|
fail(w, 400, 20001, "refreshToken required")
|
|
return
|
|
}
|
|
oldHash := sha256.Sum256([]byte(req.RefreshToken))
|
|
tx, err := a.db.BeginTx(r.Context(), &sql.TxOptions{Isolation: sql.LevelReadCommitted})
|
|
if err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "刷新登录状态失败")
|
|
return
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
var sessionID, id int64
|
|
var nickname, deviceID string
|
|
err = tx.QueryRowContext(r.Context(), `SELECT s.id,s.user_id,p.nickname,s.device_id FROM user_sessions s JOIN users u ON u.id=s.user_id JOIN user_profiles p ON p.user_id=s.user_id WHERE s.refresh_token_hash=? AND s.revoked_at IS NULL AND s.expires_at>NOW(3) AND u.status=1 AND u.deleted_at IS NULL FOR UPDATE`, oldHash[:]).Scan(&sessionID, &id, &nickname, &deviceID)
|
|
if err != nil {
|
|
fail(w, 401, 10001, "刷新令牌无效")
|
|
return
|
|
}
|
|
accessToken, err := a.token(id, "user", nickname, 30*time.Minute)
|
|
if err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "创建登录令牌失败")
|
|
return
|
|
}
|
|
newRefresh := randomToken()
|
|
newHash := sha256.Sum256([]byte(newRefresh))
|
|
result, err := tx.ExecContext(r.Context(), `UPDATE user_sessions SET refresh_token_hash=?,last_active_at=NOW(3) WHERE id=? AND refresh_token_hash=? AND revoked_at IS NULL`, newHash[:], sessionID, oldHash[:])
|
|
if err != nil {
|
|
fail(w, http.StatusInternalServerError, 50001, "刷新登录状态失败")
|
|
return
|
|
}
|
|
affected, _ := result.RowsAffected()
|
|
if affected != 1 || tx.Commit() != nil {
|
|
fail(w, http.StatusUnauthorized, 10001, "刷新令牌已被使用")
|
|
return
|
|
}
|
|
_, _ = a.db.ExecContext(r.Context(), `UPDATE user_devices SET last_active_at=NOW(3),last_ip=? WHERE user_id=? AND device_id=?`, clientIP(r), id, deviceID)
|
|
a.touchUserActivity(r.Context(), id)
|
|
reply(w, map[string]any{"accessToken": accessToken, "refreshToken": newRefresh, "expiresIn": 1800})
|
|
}
|
|
|
|
func (a *App) logout(w http.ResponseWriter, r *http.Request) {
|
|
who := current(r)
|
|
var req struct {
|
|
RefreshToken string `json:"refreshToken"`
|
|
}
|
|
if r.Body != nil && r.ContentLength != 0 {
|
|
_ = decode(r, &req)
|
|
}
|
|
if strings.TrimSpace(req.RefreshToken) != "" {
|
|
hash := sha256.Sum256([]byte(strings.TrimSpace(req.RefreshToken)))
|
|
_, _ = a.db.ExecContext(r.Context(), `UPDATE user_sessions SET revoked_at=NOW(3) WHERE user_id=? AND refresh_token_hash=? AND revoked_at IS NULL`, who.ID, hash[:])
|
|
} else {
|
|
_, _ = a.db.ExecContext(r.Context(), `UPDATE user_sessions SET revoked_at=NOW(3) WHERE user_id=? AND revoked_at IS NULL`, who.ID)
|
|
}
|
|
_, _ = a.db.ExecContext(r.Context(), `INSERT INTO user_security_controls(user_id,token_version,force_logout_at) VALUES(?,1,NOW(3)) ON DUPLICATE KEY UPDATE token_version=token_version+1,force_logout_at=VALUES(force_logout_at)`, who.ID)
|
|
a.hub.disconnect(who.ID)
|
|
reply(w, map[string]bool{"success": true})
|
|
}
|
|
|
|
func (a *App) me(w http.ResponseWriter, r *http.Request) {
|
|
profile, err := a.loadProfile(r, current(r).ID, current(r).ID)
|
|
if err != nil {
|
|
fail(w, 404, 30001, "用户不存在")
|
|
return
|
|
}
|
|
reply(w, profile)
|
|
}
|
|
|
|
func (a *App) updateProfile(w http.ResponseWriter, r *http.Request) {
|
|
var req struct {
|
|
Nickname *string `json:"nickname"`
|
|
Avatar *string `json:"avatar"`
|
|
Cover *string `json:"cover"`
|
|
Bio *string `json:"bio"`
|
|
City *string `json:"city"`
|
|
Birthday *string `json:"birthday"`
|
|
Occupation *string `json:"occupation"`
|
|
Height *int `json:"height"`
|
|
Gender *int `json:"gender"`
|
|
Education *int `json:"education"`
|
|
RelationshipStatus *int `json:"relationshipStatus"`
|
|
TagIDs *[]int64 `json:"tagIds"`
|
|
}
|
|
if err := decode(r, &req); err != nil {
|
|
fail(w, 400, 20001, err.Error())
|
|
return
|
|
}
|
|
who := current(r)
|
|
var nickname, avatar, cover, bio, city, occupation string
|
|
var birthday sql.NullString
|
|
var height sql.NullInt64
|
|
var gender, education, relationship int
|
|
if err := a.db.QueryRowContext(r.Context(), `SELECT nickname,avatar_url,cover_url,bio,city_name,occupation,DATE_FORMAT(birthday,'%Y-%m-%d'),height_cm,gender,education,relationship_status FROM user_profiles WHERE user_id=?`, who.ID).Scan(&nickname, &avatar, &cover, &bio, &city, &occupation, &birthday, &height, &gender, &education, &relationship); err != nil {
|
|
fail(w, 500, 50001, "读取资料失败")
|
|
return
|
|
}
|
|
if req.Nickname != nil {
|
|
nickname = strings.TrimSpace(*req.Nickname)
|
|
}
|
|
if req.Avatar != nil {
|
|
avatar = strings.TrimSpace(*req.Avatar)
|
|
}
|
|
if req.Cover != nil {
|
|
cover = strings.TrimSpace(*req.Cover)
|
|
}
|
|
if req.Bio != nil {
|
|
bio = strings.TrimSpace(*req.Bio)
|
|
}
|
|
if req.City != nil {
|
|
city = strings.TrimSpace(*req.City)
|
|
}
|
|
if req.Occupation != nil {
|
|
occupation = strings.TrimSpace(*req.Occupation)
|
|
}
|
|
if req.Height != nil {
|
|
if *req.Height < 0 || *req.Height > 260 {
|
|
fail(w, 400, 20001, "身高范围无效")
|
|
return
|
|
}
|
|
height = sql.NullInt64{Int64: int64(*req.Height), Valid: *req.Height > 0}
|
|
}
|
|
if req.Gender != nil {
|
|
if *req.Gender < 0 || *req.Gender > 2 {
|
|
fail(w, 400, 20001, "性别选项无效")
|
|
return
|
|
}
|
|
gender = *req.Gender
|
|
}
|
|
if req.Education != nil {
|
|
if *req.Education < 0 || *req.Education > 10 {
|
|
fail(w, 400, 20001, "学历选项无效")
|
|
return
|
|
}
|
|
education = *req.Education
|
|
}
|
|
if req.RelationshipStatus != nil {
|
|
if *req.RelationshipStatus < 0 || *req.RelationshipStatus > 10 {
|
|
fail(w, 400, 20001, "情感状态无效")
|
|
return
|
|
}
|
|
relationship = *req.RelationshipStatus
|
|
}
|
|
if nickname == "" || len([]rune(nickname)) > 50 || len([]rune(bio)) > 500 || len([]rune(city)) > 50 || len([]rune(occupation)) > 100 {
|
|
fail(w, 400, 20001, "资料内容长度无效")
|
|
return
|
|
}
|
|
if req.Birthday != nil {
|
|
value := strings.TrimSpace(*req.Birthday)
|
|
if value == "" {
|
|
birthday = sql.NullString{}
|
|
} else {
|
|
parsed, parseErr := time.Parse("2006-01-02", value)
|
|
if parseErr != nil || parsed.After(time.Now().AddDate(-18, 0, 0)) || parsed.Before(time.Now().AddDate(-100, 0, 0)) {
|
|
fail(w, 400, 20001, "仅支持 18-100 周岁的生日日期")
|
|
return
|
|
}
|
|
birthday = sql.NullString{String: value, Valid: true}
|
|
}
|
|
}
|
|
tx, err := a.db.BeginTx(r.Context(), nil)
|
|
if err != nil {
|
|
fail(w, 500, 50001, "保存失败")
|
|
return
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
_, err = tx.ExecContext(r.Context(), `UPDATE user_profiles SET nickname=?,avatar_url=?,cover_url=?,bio=?,city_name=?,occupation=?,birthday=?,height_cm=?,gender=?,education=?,relationship_status=?,profile_score=GREATEST(profile_score,80) WHERE user_id=?`, nickname, avatar, cover, bio, city, occupation, birthday, height, gender, education, relationship, who.ID)
|
|
if err == nil && req.TagIDs != nil {
|
|
if len(*req.TagIDs) > 12 {
|
|
fail(w, 400, 20001, "最多选择 12 个标签")
|
|
return
|
|
}
|
|
_, err = tx.ExecContext(r.Context(), `DELETE FROM user_tags WHERE user_id=?`, who.ID)
|
|
for _, tagID := range *req.TagIDs {
|
|
if err != nil {
|
|
break
|
|
}
|
|
result, insertErr := tx.ExecContext(r.Context(), `INSERT INTO user_tags(user_id,tag_id) SELECT ?,id FROM tags WHERE id=? AND status=1`, who.ID, tagID)
|
|
err = insertErr
|
|
if err == nil {
|
|
affected, _ := result.RowsAffected()
|
|
if affected == 0 {
|
|
err = fmt.Errorf("标签不存在")
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if err != nil || tx.Commit() != nil {
|
|
fail(w, 500, 50001, "保存失败")
|
|
return
|
|
}
|
|
a.me(w, r)
|
|
}
|
|
|
|
func nullableString(v sql.NullString) string {
|
|
if v.Valid {
|
|
return v.String
|
|
}
|
|
return ""
|
|
}
|