This commit is contained in:
Your Name
2026-08-27 14:04:28 +08:00
parent f7720831be
commit 334890171e
3016 changed files with 263403 additions and 27971 deletions
+344
View File
@@ -0,0 +1,344 @@
package app
import (
"context"
"database/sql"
"errors"
"net/http"
"net/url"
"strings"
"time"
)
type userOAuthLoginCode struct {
Provider string
Subject string
Email string
DisplayName string
AvatarURL string
UserID sql.NullInt64
}
func (a *App) userOAuthFrontendURL(ctx context.Context) (string, error) {
raw := strings.TrimSpace(a.configPlain(ctx, "oauth.user.frontend_callback_url", ""))
if raw == "" {
return "", errors.New("客户端第三方登录结果页未配置")
}
if err := a.validateAdminOAuthRedirectURL(raw); err != nil {
return "", err
}
return raw, nil
}
func (a *App) enabledUserOAuthProviders(ctx context.Context) ([]adminOAuthProvider, error) {
providers := make([]adminOAuthProvider, 0, len(adminOAuthProviderNames))
for _, code := range []string{"wechat", "qq", "github", "google"} {
if !a.configBool(ctx, "oauth.user."+code+".enabled", false) {
continue
}
provider, err := a.adminOAuthProvider(ctx, code)
if err != nil {
return nil, err
}
providers = append(providers, provider)
}
return providers, nil
}
// oauthCallback lets one provider callback URL safely serve both the admin
// console and the uni-app H5 client. The random state value selects the
// audience; it is never accepted by both state tables.
func (a *App) oauthCallback(w http.ResponseWriter, r *http.Request) {
state := strings.TrimSpace(r.URL.Query().Get("state"))
if state != "" {
var exists int
if a.db.QueryRowContext(r.Context(), `SELECT 1 FROM user_oauth_states WHERE state_hash=?`, oauthHash(state)).Scan(&exists) == nil {
a.userOAuthCallback(w, r)
return
}
if a.db.QueryRowContext(r.Context(), `SELECT 1 FROM admin_oauth_states WHERE state_hash=?`, oauthHash(state)).Scan(&exists) == nil {
a.adminOAuthCallback(w, r)
return
}
}
if strings.HasPrefix(r.URL.Path, "/api/") {
a.userOAuthCallback(w, r)
return
}
a.adminOAuthCallback(w, r)
}
func (a *App) userOAuthProviders(w http.ResponseWriter, r *http.Request) {
items := make([]map[string]string, 0, len(adminOAuthProviderNames))
for _, code := range []string{"wechat", "qq", "github", "google"} {
if !a.configBool(r.Context(), "oauth.user."+code+".enabled", false) {
continue
}
provider, err := a.adminOAuthProvider(r.Context(), code)
if err != nil {
continue
}
items = append(items, map[string]string{"code": provider.Code, "name": provider.Name})
}
reply(w, map[string]any{"items": items})
}
func (a *App) userOAuthStart(w http.ResponseWriter, r *http.Request) {
var req struct {
Provider string `json:"provider"`
}
if decode(r, &req) != nil {
fail(w, http.StatusBadRequest, 20001, "请选择第三方登录渠道")
return
}
req.Provider = strings.ToLower(strings.TrimSpace(req.Provider))
if !a.rateLimit(w, r, "user_oauth_start", clientIP(r), 30, 10*time.Minute) {
return
}
if !a.configBool(r.Context(), "oauth.user."+req.Provider+".enabled", false) {
fail(w, http.StatusBadRequest, 20001, "该客户端登录方式未启用")
return
}
provider, err := a.adminOAuthProvider(r.Context(), req.Provider)
if err != nil {
fail(w, http.StatusBadRequest, 20001, "该登录方式配置不完整")
return
}
if _, err = a.userOAuthFrontendURL(r.Context()); err != nil {
fail(w, http.StatusBadRequest, 20001, "客户端登录结果页配置不完整")
return
}
state := randomToken()
verifier := ""
if provider.Code == "github" || provider.Code == "google" {
verifier = randomToken() + randomToken()
}
_, err = a.db.ExecContext(r.Context(), `INSERT INTO user_oauth_states(state_hash,provider,code_verifier,expires_at) VALUES(?,?,?,?)`, oauthHash(state), provider.Code, verifier, time.Now().Add(adminOAuthStateTTL))
if err != nil {
fail(w, http.StatusInternalServerError, 50001, "创建第三方登录请求失败")
return
}
a.cleanupUserOAuthRecords(r.Context())
authorizationURL, _ := url.Parse(provider.AuthorizationURL)
query := authorizationURL.Query()
if provider.Code == "wechat" {
query.Set("appid", provider.ClientID)
} else {
query.Set("client_id", provider.ClientID)
}
query.Set("redirect_uri", provider.RedirectURI)
query.Set("response_type", "code")
query.Set("scope", provider.Scope)
query.Set("state", state)
if verifier != "" {
query.Set("code_challenge", pkceChallenge(verifier))
query.Set("code_challenge_method", "S256")
}
authorizationURL.RawQuery = query.Encode()
if provider.Code == "wechat" {
authorizationURL.Fragment = "wechat_redirect"
}
reply(w, map[string]string{"authorizationUrl": authorizationURL.String(), "provider": provider.Code})
}
func (a *App) userOAuthCallback(w http.ResponseWriter, r *http.Request) {
frontendURL, err := a.userOAuthFrontendURL(r.Context())
if err != nil {
fail(w, http.StatusServiceUnavailable, 50001, "客户端第三方登录回调未配置")
return
}
state := strings.TrimSpace(r.URL.Query().Get("state"))
if state == "" {
a.redirectUserOAuthResult(w, r, frontendURL, "", "登录状态无效或已过期")
return
}
var providerCode, verifier string
err = a.db.QueryRowContext(r.Context(), `SELECT provider,code_verifier FROM user_oauth_states WHERE state_hash=? AND used_at IS NULL AND expires_at>NOW(3)`, oauthHash(state)).Scan(&providerCode, &verifier)
if err != nil {
a.redirectUserOAuthResult(w, r, frontendURL, "", "登录状态无效或已过期")
return
}
result, err := a.db.ExecContext(r.Context(), `UPDATE user_oauth_states SET used_at=NOW(3) WHERE state_hash=? AND used_at IS NULL AND expires_at>NOW(3)`, oauthHash(state))
if err != nil {
a.redirectUserOAuthResult(w, r, frontendURL, "", "第三方登录处理失败")
return
}
affected, _ := result.RowsAffected()
if affected != 1 {
a.redirectUserOAuthResult(w, r, frontendURL, "", "登录状态已被使用")
return
}
if strings.TrimSpace(r.URL.Query().Get("error")) != "" {
a.redirectUserOAuthResult(w, r, frontendURL, "", "第三方授权已取消或失败")
return
}
code := strings.TrimSpace(r.URL.Query().Get("code"))
if code == "" {
a.redirectUserOAuthResult(w, r, frontendURL, "", "第三方平台未返回授权码")
return
}
if !a.configBool(r.Context(), "oauth.user."+providerCode+".enabled", false) {
a.redirectUserOAuthResult(w, r, frontendURL, "", "该客户端登录方式已停用")
return
}
provider, err := a.adminOAuthProvider(r.Context(), providerCode)
if err != nil {
a.redirectUserOAuthResult(w, r, frontendURL, "", "该登录方式配置不可用")
return
}
identity, err := a.fetchAdminOAuthIdentity(r.Context(), provider, code, verifier)
if err != nil {
a.redirectUserOAuthResult(w, r, frontendURL, "", "获取第三方账号信息失败")
return
}
var userID sql.NullInt64
_ = a.db.QueryRowContext(r.Context(), `SELECT user_id FROM user_oauth_identities WHERE provider=? AND subject=?`, provider.Code, identity.Subject).Scan(&userID)
loginCode := randomToken()
_, err = a.db.ExecContext(r.Context(), `INSERT INTO user_oauth_login_codes(code_hash,provider,subject,email,display_name,avatar_url,user_id,expires_at) VALUES(?,?,?,?,?,?,?,?)`, oauthHash(loginCode), provider.Code, identity.Subject, identity.Email, identity.DisplayName, identity.AvatarURL, userID, time.Now().Add(adminOAuthCodeTTL))
if err != nil {
a.redirectUserOAuthResult(w, r, frontendURL, "", "创建登录凭证失败")
return
}
a.redirectUserOAuthResult(w, r, frontendURL, loginCode, "")
}
func (a *App) redirectUserOAuthResult(w http.ResponseWriter, r *http.Request, frontendURL, code, message string) {
target, err := url.Parse(frontendURL)
if err != nil {
fail(w, http.StatusInternalServerError, 50001, "客户端登录结果页地址无效")
return
}
query := target.Query()
if code != "" {
query.Set("oauthCode", code)
} else {
query.Set("oauthError", message)
}
target.RawQuery = query.Encode()
http.Redirect(w, r, target.String(), http.StatusFound)
}
func (a *App) userOAuthExchange(w http.ResponseWriter, r *http.Request) {
var req struct {
Code string `json:"code"`
DeviceID string `json:"deviceId"`
}
if decode(r, &req) != nil || strings.TrimSpace(req.Code) == "" {
fail(w, http.StatusBadRequest, 20001, "第三方登录凭证无效")
return
}
if !a.rateLimit(w, r, "user_oauth_exchange", clientIP(r), 20, 10*time.Minute) {
return
}
loginCode, err := a.readUserOAuthLoginCode(r.Context(), strings.TrimSpace(req.Code))
if err != nil {
fail(w, http.StatusBadRequest, 20001, "第三方登录凭证无效或已过期")
return
}
if !loginCode.UserID.Valid {
reply(w, map[string]any{
"requiresLink": true,
"provider": loginCode.Provider,
"providerName": adminOAuthProviderNames[loginCode.Provider],
"displayName": loginCode.DisplayName,
"avatarUrl": loginCode.AvatarURL,
})
return
}
userID, nickname, err := a.consumeUserOAuthCode(r.Context(), strings.TrimSpace(req.Code), loginCode.UserID.Int64)
if err != nil {
fail(w, http.StatusUnauthorized, 10001, err.Error())
return
}
a.finishLogin(w, r, userID, nickname, req.DeviceID)
}
func (a *App) userOAuthLink(w http.ResponseWriter, r *http.Request) {
var req struct {
Code string `json:"code"`
Phone string `json:"phone"`
SMSCode string `json:"smsCode"`
DeviceID string `json:"deviceId"`
}
if decode(r, &req) != nil || strings.TrimSpace(req.Code) == "" || !validPhone(req.Phone) || len(req.SMSCode) != 6 {
fail(w, http.StatusBadRequest, 20001, "请输入已注册手机号和正确的短信验证码")
return
}
phone := strings.TrimSpace(req.Phone)
if !a.rateLimit(w, r, "user_oauth_link_ip", clientIP(r), 10, 15*time.Minute) || !a.rateLimit(w, r, "user_oauth_link_phone", phone, 10, 15*time.Minute) {
return
}
loginCode, err := a.readUserOAuthLoginCode(r.Context(), strings.TrimSpace(req.Code))
if err != nil || loginCode.UserID.Valid {
fail(w, http.StatusBadRequest, 20001, "第三方登录凭证无效、已绑定或已过期")
return
}
var userID int64
var nickname string
var status int
err = a.db.QueryRowContext(r.Context(), `SELECT u.id,p.nickname,u.status FROM users u JOIN user_profiles p ON p.user_id=u.id WHERE u.phone_hash=? AND u.deleted_at IS NULL`, phoneHash(phone)).Scan(&userID, &nickname, &status)
if err != nil || status != 1 {
fail(w, http.StatusBadRequest, 20001, "手机号未注册或账号当前不可用")
return
}
if !a.consumeSMSCode(r, phone, "login", req.SMSCode) {
fail(w, http.StatusBadRequest, 20001, "验证码错误或已过期")
return
}
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() }()
result, err := tx.ExecContext(r.Context(), `UPDATE user_oauth_login_codes SET used_at=NOW(3),user_id=? WHERE code_hash=? AND user_id IS NULL AND used_at IS NULL AND expires_at>NOW(3)`, userID, oauthHash(strings.TrimSpace(req.Code)))
if err != nil {
fail(w, http.StatusInternalServerError, 50001, "绑定第三方账号失败")
return
}
affected, _ := result.RowsAffected()
if affected != 1 {
fail(w, http.StatusBadRequest, 20001, "第三方登录凭证已被使用")
return
}
_, err = tx.ExecContext(r.Context(), `INSERT INTO user_oauth_identities(provider,subject,user_id,email,display_name,avatar_url,last_login_at) VALUES(?,?,?,?,?,?,NOW(3))`, loginCode.Provider, loginCode.Subject, userID, loginCode.Email, loginCode.DisplayName, loginCode.AvatarURL)
if err != nil {
fail(w, http.StatusConflict, 20001, "该第三方账号或手机号已绑定此渠道")
return
}
if err = tx.Commit(); err != nil {
fail(w, http.StatusInternalServerError, 50001, "绑定第三方账号失败")
return
}
a.finishLogin(w, r, userID, nickname, req.DeviceID)
}
func (a *App) readUserOAuthLoginCode(ctx context.Context, code string) (userOAuthLoginCode, error) {
var result userOAuthLoginCode
err := a.db.QueryRowContext(ctx, `SELECT provider,subject,email,display_name,avatar_url,user_id FROM user_oauth_login_codes WHERE code_hash=? AND used_at IS NULL AND expires_at>NOW(3)`, oauthHash(code)).Scan(&result.Provider, &result.Subject, &result.Email, &result.DisplayName, &result.AvatarURL, &result.UserID)
return result, err
}
func (a *App) consumeUserOAuthCode(ctx context.Context, code string, userID int64) (int64, string, error) {
var nickname string
var status int
if err := a.db.QueryRowContext(ctx, `SELECT p.nickname,u.status FROM users u JOIN user_profiles p ON p.user_id=u.id WHERE u.id=? AND u.deleted_at IS NULL`, userID).Scan(&nickname, &status); err != nil || status != 1 {
return 0, "", errors.New("账号不存在或当前不可用")
}
result, err := a.db.ExecContext(ctx, `UPDATE user_oauth_login_codes SET used_at=NOW(3) WHERE code_hash=? AND user_id=? AND used_at IS NULL AND expires_at>NOW(3)`, oauthHash(code), userID)
if err != nil {
return 0, "", errors.New("第三方登录处理失败")
}
affected, _ := result.RowsAffected()
if affected != 1 {
return 0, "", errors.New("第三方登录凭证无效或已使用")
}
_, _ = a.db.ExecContext(ctx, `UPDATE user_oauth_identities SET last_login_at=NOW(3) WHERE user_id=?`, userID)
return userID, nickname, nil
}
func (a *App) cleanupUserOAuthRecords(ctx context.Context) {
_, _ = a.db.ExecContext(ctx, `DELETE FROM user_oauth_states WHERE expires_at<DATE_SUB(NOW(3),INTERVAL 1 DAY) LIMIT 500`)
_, _ = a.db.ExecContext(ctx, `DELETE FROM user_oauth_login_codes WHERE expires_at<DATE_SUB(NOW(3),INTERVAL 1 DAY) LIMIT 500`)
}