批量更新
This commit is contained in:
@@ -0,0 +1,413 @@
|
||||
// Package idp 第三方身份源(微信/钉钉/飞书/QQ/GitHub/Google)统一适配层。
|
||||
//
|
||||
// 各平台 OAuth2 流程基本一致,差异只在端点地址、参数名与用户字段,
|
||||
// 因此用一份通用实现 + 预设配置覆盖,新增平台只需加一条 preset。
|
||||
package idp
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"server/models"
|
||||
)
|
||||
|
||||
// UserInfo 第三方返回的标准化用户信息
|
||||
type UserInfo struct {
|
||||
OpenID string // 该平台内的唯一 ID(必填)
|
||||
UnionID string // 跨平台唯一 ID(微信/钉钉有,GitHub/Google 用 OpenID 代替)
|
||||
Nickname string
|
||||
Avatar string
|
||||
Raw string // 原始 JSON,便于排查
|
||||
}
|
||||
|
||||
// Config 一个第三方身份源的配置
|
||||
type Config struct {
|
||||
Provider string
|
||||
AuthURL string
|
||||
TokenURL string
|
||||
UserinfoURL string
|
||||
Scopes []string
|
||||
AppID string
|
||||
AppSecret string
|
||||
ProxyURL string // GitHub/Google 在国内服务器需要代理时填写
|
||||
|
||||
// 参数风格:微信/QQ 用 appid+secret,其余用 client_id+client_secret
|
||||
UseAppIDStyle bool
|
||||
// 用户信息接口风格:query=拼在 URL 上(微信/QQ),bearer=放 Authorization 头
|
||||
UserinfoStyle string
|
||||
|
||||
// 用户信息字段映射(一级 JSON 字段)
|
||||
OpenIDField string
|
||||
UnionIDField string
|
||||
NicknameField string
|
||||
AvatarField string
|
||||
}
|
||||
|
||||
// presets 各平台预设(端点与字段映射)
|
||||
var presets = map[string]Config{
|
||||
models.IdPWechat: {
|
||||
Provider: models.IdPWechat,
|
||||
AuthURL: "https://open.weixin.qq.com/connect/oauth2/authorize",
|
||||
TokenURL: "https://api.weixin.qq.com/sns/oauth2/access_token",
|
||||
UserinfoURL: "https://api.weixin.qq.com/sns/userinfo",
|
||||
Scopes: []string{"snsapi_userinfo"},
|
||||
UseAppIDStyle: true,
|
||||
UserinfoStyle: "query",
|
||||
OpenIDField: "openid",
|
||||
UnionIDField: "unionid",
|
||||
NicknameField: "nickname",
|
||||
AvatarField: "headimgurl",
|
||||
},
|
||||
models.IdPQQ: {
|
||||
Provider: models.IdPQQ,
|
||||
AuthURL: "https://graph.qq.com/oauth2.0/authorize",
|
||||
TokenURL: "https://graph.qq.com/oauth2.0/token",
|
||||
UserinfoURL: "https://graph.qq.com/user/get_user_info",
|
||||
Scopes: []string{"get_user_info"},
|
||||
UseAppIDStyle: true,
|
||||
UserinfoStyle: "query",
|
||||
OpenIDField: "openid",
|
||||
NicknameField: "nickname",
|
||||
AvatarField: "figureurl_qq_2",
|
||||
},
|
||||
models.IdPGitHub: {
|
||||
Provider: models.IdPGitHub,
|
||||
AuthURL: "https://github.com/login/oauth/authorize",
|
||||
TokenURL: "https://github.com/login/oauth/access_token",
|
||||
UserinfoURL: "https://api.github.com/user",
|
||||
Scopes: []string{"read:user"},
|
||||
UserinfoStyle: "bearer",
|
||||
OpenIDField: "id",
|
||||
NicknameField: "login",
|
||||
AvatarField: "avatar_url",
|
||||
},
|
||||
models.IdPGoogle: {
|
||||
Provider: models.IdPGoogle,
|
||||
AuthURL: "https://accounts.google.com/o/oauth2/v2/auth",
|
||||
TokenURL: "https://oauth2.googleapis.com/token",
|
||||
UserinfoURL: "https://openidconnect.googleapis.com/v1/userinfo",
|
||||
Scopes: []string{"openid", "profile"},
|
||||
UserinfoStyle: "bearer",
|
||||
OpenIDField: "sub",
|
||||
NicknameField: "name",
|
||||
AvatarField: "picture",
|
||||
},
|
||||
models.IdPDingTalk: {
|
||||
Provider: models.IdPDingTalk,
|
||||
AuthURL: "https://login.dingtalk.com/oauth2/auth",
|
||||
TokenURL: "https://api.dingtalk.com/v1.0/oauth2/userAccessToken",
|
||||
UserinfoURL: "https://api.dingtalk.com/v1.0/contact/users/me",
|
||||
Scopes: []string{"openid", "profile"},
|
||||
UserinfoStyle: "bearer",
|
||||
OpenIDField: "unionId",
|
||||
UnionIDField: "unionId",
|
||||
NicknameField: "nick",
|
||||
AvatarField: "avatarUrl",
|
||||
},
|
||||
models.IdPFeishu: {
|
||||
Provider: models.IdPFeishu,
|
||||
AuthURL: "https://open.feishu.cn/open-apis/authen/v1/authorize",
|
||||
TokenURL: "https://open.feishu.cn/open-apis/authen/v2/oauth/token",
|
||||
UserinfoURL: "https://open.feishu.cn/open-apis/authen/v1/user_info",
|
||||
Scopes: []string{"contact:user.base:readonly"},
|
||||
UserinfoStyle: "bearer",
|
||||
OpenIDField: "open_id",
|
||||
UnionIDField: "union_id",
|
||||
NicknameField: "name",
|
||||
AvatarField: "avatar_url",
|
||||
},
|
||||
}
|
||||
|
||||
// Supported 返回支持的第三方平台列表
|
||||
func Supported() []string {
|
||||
return []string{
|
||||
models.IdPWechat, models.IdPDingTalk, models.IdPFeishu,
|
||||
models.IdPQQ, models.IdPGitHub, models.IdPGoogle,
|
||||
}
|
||||
}
|
||||
|
||||
// IsSupported 是否为已知平台
|
||||
func IsSupported(provider string) bool {
|
||||
_, ok := presets[provider]
|
||||
return ok
|
||||
}
|
||||
|
||||
// LoadConfig 读取身份源配置:
|
||||
// - tid > 0 时优先取租户自带身份源(yz_auth_tenant_idp)
|
||||
// - 取不到或 tid=0 时取平台全局配置(tid=0 的记录)
|
||||
//
|
||||
// 未配置任何记录时返回 false,调用方应提示「该登录方式未开通」。
|
||||
func LoadConfig(provider string, tid uint64) (Config, bool) {
|
||||
base, ok := presets[provider]
|
||||
if !ok {
|
||||
return Config{}, false
|
||||
}
|
||||
|
||||
var row models.AuthTenantIdp
|
||||
qs := models.Orm.QueryTable(new(models.AuthTenantIdp)).
|
||||
Filter("provider", provider).
|
||||
Filter("status", 1)
|
||||
|
||||
if tid > 0 {
|
||||
if err := qs.Filter("tid", tid).One(&row); err == nil {
|
||||
return applyDBConfig(base, &row), true
|
||||
}
|
||||
}
|
||||
// 回落到平台全局配置(tid=0)
|
||||
if err := qs.Filter("tid", 0).One(&row); err == nil {
|
||||
return applyDBConfig(base, &row), true
|
||||
}
|
||||
return base, false
|
||||
}
|
||||
|
||||
func applyDBConfig(base Config, row *models.AuthTenantIdp) Config {
|
||||
if row.AppID != nil && strings.TrimSpace(*row.AppID) != "" {
|
||||
base.AppID = strings.TrimSpace(*row.AppID)
|
||||
}
|
||||
if row.AppSecret != nil && strings.TrimSpace(*row.AppSecret) != "" {
|
||||
base.AppSecret = strings.TrimSpace(*row.AppSecret)
|
||||
}
|
||||
if row.ProxyURL != nil && strings.TrimSpace(*row.ProxyURL) != "" {
|
||||
base.ProxyURL = strings.TrimSpace(*row.ProxyURL)
|
||||
}
|
||||
if row.AuthURL != nil && strings.TrimSpace(*row.AuthURL) != "" {
|
||||
base.AuthURL = strings.TrimSpace(*row.AuthURL)
|
||||
}
|
||||
if row.TokenURL != nil && strings.TrimSpace(*row.TokenURL) != "" {
|
||||
base.TokenURL = strings.TrimSpace(*row.TokenURL)
|
||||
}
|
||||
if row.UserinfoURL != nil && strings.TrimSpace(*row.UserinfoURL) != "" {
|
||||
base.UserinfoURL = strings.TrimSpace(*row.UserinfoURL)
|
||||
}
|
||||
if row.Scopes != nil && strings.TrimSpace(*row.Scopes) != "" {
|
||||
base.Scopes = splitList(*row.Scopes)
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
func splitList(s string) []string {
|
||||
parts := strings.Split(s, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
if v := strings.TrimSpace(p); v != "" {
|
||||
out = append(out, v)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// BuildAuthURL 生成跳转到第三方授权页的地址
|
||||
func BuildAuthURL(cfg Config, redirectURI, state string) string {
|
||||
v := url.Values{}
|
||||
if cfg.UseAppIDStyle {
|
||||
v.Set("appid", cfg.AppID)
|
||||
} else {
|
||||
v.Set("client_id", cfg.AppID)
|
||||
}
|
||||
v.Set("redirect_uri", redirectURI)
|
||||
v.Set("response_type", "code")
|
||||
v.Set("scope", strings.Join(cfg.Scopes, ","))
|
||||
v.Set("state", state)
|
||||
|
||||
addr := cfg.AuthURL + "?" + v.Encode()
|
||||
if cfg.Provider == models.IdPWechat {
|
||||
// 微信要求在 hash 后带 #wechat_redirect
|
||||
return addr + "#wechat_redirect"
|
||||
}
|
||||
return addr
|
||||
}
|
||||
|
||||
// Exchange 用授权码换取用户信息
|
||||
func Exchange(cfg Config, code string) (*UserInfo, error) {
|
||||
if cfg.AppID == "" || cfg.AppSecret == "" {
|
||||
return nil, errors.New("该登录方式尚未配置,请联系管理员")
|
||||
}
|
||||
|
||||
// 1. code 换 access_token
|
||||
v := url.Values{}
|
||||
if cfg.UseAppIDStyle {
|
||||
v.Set("appid", cfg.AppID)
|
||||
v.Set("secret", cfg.AppSecret)
|
||||
} else {
|
||||
v.Set("client_id", cfg.AppID)
|
||||
v.Set("client_secret", cfg.AppSecret)
|
||||
}
|
||||
v.Set("code", code)
|
||||
v.Set("grant_type", "authorization_code")
|
||||
|
||||
tokenBody, err := post(cfg, cfg.TokenURL, v, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tokenMap := parseMap(tokenBody)
|
||||
accessToken := firstStr(tokenMap, "access_token")
|
||||
if accessToken == "" {
|
||||
return nil, fmt.Errorf("换取令牌失败: %s", truncate(tokenBody, 200))
|
||||
}
|
||||
openID := firstStr(tokenMap, "openid", "unionId", "open_id", "sub", "id")
|
||||
|
||||
// 2. 取用户信息
|
||||
var infoMap map[string]interface{}
|
||||
if cfg.UserinfoStyle == "query" {
|
||||
q := url.Values{}
|
||||
q.Set("access_token", accessToken)
|
||||
if openID != "" {
|
||||
q.Set("openid", openID)
|
||||
}
|
||||
// QQ 需要额外带 oauth_consumer_key
|
||||
if cfg.Provider == models.IdPQQ {
|
||||
q.Set("oauth_consumer_key", cfg.AppID)
|
||||
}
|
||||
body, err := get(cfg, cfg.UserinfoURL+"?"+q.Encode(), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
infoMap = parseMap(body)
|
||||
} else {
|
||||
body, err := get(cfg, cfg.UserinfoURL, map[string]string{
|
||||
"Authorization": "Bearer " + accessToken,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
infoMap = parseMap(body)
|
||||
// 飞书把数据包在 data 里
|
||||
if d, ok := infoMap["data"].(map[string]interface{}); ok {
|
||||
infoMap = d
|
||||
}
|
||||
}
|
||||
|
||||
info := &UserInfo{Raw: truncate(mustJSON(infoMap), 2000)}
|
||||
info.OpenID = strField(infoMap, cfg.OpenIDField)
|
||||
if cfg.UnionIDField != "" {
|
||||
info.UnionID = strField(infoMap, cfg.UnionIDField)
|
||||
}
|
||||
info.Nickname = strField(infoMap, cfg.NicknameField)
|
||||
info.Avatar = strField(infoMap, cfg.AvatarField)
|
||||
|
||||
if info.OpenID == "" && openID != "" {
|
||||
info.OpenID = openID
|
||||
}
|
||||
if info.OpenID == "" {
|
||||
return nil, errors.New("未能获取第三方账号标识")
|
||||
}
|
||||
if info.UnionID == "" {
|
||||
info.UnionID = info.OpenID
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- HTTP
|
||||
|
||||
func client(cfg Config) *http.Client {
|
||||
c := &http.Client{Timeout: 15 * time.Second}
|
||||
if cfg.ProxyURL != "" {
|
||||
if u, err := url.Parse(cfg.ProxyURL); err == nil {
|
||||
c.Transport = &http.Transport{Proxy: http.ProxyURL(u)}
|
||||
}
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func post(cfg Config, target string, form url.Values, headers map[string]string) (string, error) {
|
||||
req, err := http.NewRequest("POST", target, strings.NewReader(form.Encode()))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
if cfg.Provider == models.IdPGitHub {
|
||||
req.Header.Set("Accept", "application/json")
|
||||
}
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
resp, err := client(cfg).Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求第三方失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode >= 300 {
|
||||
return "", fmt.Errorf("第三方返回异常(%d): %s", resp.StatusCode, truncate(string(body), 200))
|
||||
}
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
func get(cfg Config, target string, headers map[string]string) (string, error) {
|
||||
req, err := http.NewRequest("GET", target, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
resp, err := client(cfg).Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求第三方失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode >= 300 {
|
||||
return "", fmt.Errorf("第三方返回异常(%d): %s", resp.StatusCode, truncate(string(body), 200))
|
||||
}
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 解析工具
|
||||
|
||||
// parseMap 兼容 JSON 与 querystring 两种响应(微信/QQ 早期接口返回 form 格式)
|
||||
func parseMap(body string) map[string]interface{} {
|
||||
out := map[string]interface{}{}
|
||||
if err := json.Unmarshal([]byte(body), &out); err == nil {
|
||||
return out
|
||||
}
|
||||
if values, err := url.ParseQuery(body); err == nil {
|
||||
for k, v := range values {
|
||||
if len(v) > 0 {
|
||||
out[k] = v[0]
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func firstStr(m map[string]interface{}, keys ...string) string {
|
||||
for _, k := range keys {
|
||||
if v, ok := m[k]; ok {
|
||||
switch t := v.(type) {
|
||||
case string:
|
||||
if t != "" {
|
||||
return t
|
||||
}
|
||||
case float64:
|
||||
return fmt.Sprintf("%.0f", t)
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func strField(m map[string]interface{}, field string) string {
|
||||
if field == "" {
|
||||
return ""
|
||||
}
|
||||
return firstStr(m, field)
|
||||
}
|
||||
|
||||
func mustJSON(m map[string]interface{}) string {
|
||||
b, _ := json.Marshal(m)
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func truncate(s string, n int) string {
|
||||
if len(s) <= n {
|
||||
return s
|
||||
}
|
||||
return s[:n]
|
||||
}
|
||||
Reference in New Issue
Block a user