Files
yunzerwebsiteallinone/go/controllers/auth/oidc.go
T
2026-09-25 00:27:15 +08:00

605 lines
17 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Package auth 统一认证中心(UAC)控制器:api.yunzer.cn/auth
//
// 实现 OIDC 1.0(基于 OAuth 2.1 + PKCE)标准端点,
// 以后每开发一个新软件,只需在 yz_auth_client 注册一条即可接入。
package auth
import (
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"net/url"
"strings"
"time"
authsvc "server/services/auth"
"server/models"
"server/pkg/jwtutil"
beego "github.com/beego/beego/v2/server/web"
)
// 认证中心会话 Cookie(仅作用于认证中心域名,用于 authorize 阶段识别登录态)
const (
sessionCookieName = "yz_sid"
sessionCookieTTL = 7200
)
// grant_type 常量
const (
GrantAuthCode = "authorization_code"
GrantRefresh = "refresh_token"
ResponseTypeCode = "code"
)
// AuthOidcController OIDC 标准端点
type AuthOidcController struct {
beego.Controller
}
func (c *AuthOidcController) serveJSON(data map[string]interface{}) {
c.Data["json"] = data
_ = c.ServeJSON()
}
func (c *AuthOidcController) fail(status int, msg string) {
c.Ctx.Output.SetStatus(status)
c.serveJSON(map[string]interface{}{"error": msg})
}
// Discovery OIDC 发现文档
// GET /auth/.well-known/openid-configuration
func (c *AuthOidcController) Discovery() {
issuer := jwtutil.Issuer()
if issuer == "" {
issuer = fmt.Sprintf("https://%s/auth", c.Ctx.Request.Host)
}
c.Data["json"] = map[string]interface{}{
"issuer": issuer,
"authorization_endpoint": issuer + "/authorize",
"token_endpoint": issuer + "/token",
"userinfo_endpoint": issuer + "/userinfo",
"introspection_endpoint": issuer + "/introspect",
"revocation_endpoint": issuer + "/revoke",
"end_session_endpoint": issuer + "/logout",
"jwks_uri": issuer + "/jwks.json",
"response_types_supported": []string{"code"},
"grant_types_supported": []string{GrantAuthCode, GrantRefresh},
"subject_types_supported": []string{"public"},
"id_token_signing_alg_values_supported": []string{jwtutil.AlgRS256, jwtutil.AlgHS256},
"code_challenge_methods_supported": []string{"S256"},
"scopes_supported": []string{"openid", "profile", "tenant"},
}
_ = c.ServeJSON()
}
// JWKS 公钥集合(各应用本地验签用)
// GET /auth/jwks.json
func (c *AuthOidcController) JWKS() {
keys := jwtutil.JWKS()
if keys == nil {
keys = []jwtutil.JWK{}
}
c.Data["json"] = map[string]interface{}{"keys": keys}
_ = c.ServeJSON()
}
// Authorize 授权端点
// GET /auth/authorize?client_id=&redirect_uri=&response_type=code&scope=&state=&code_challenge=&code_challenge_method=S256
//
// 未登录时重定向到统一登录页,登录后再回到本端点完成授权。
func (c *AuthOidcController) Authorize() {
clientID := strings.TrimSpace(c.GetString("client_id"))
redirectURI := strings.TrimSpace(c.GetString("redirect_uri"))
responseType := strings.TrimSpace(c.GetString("response_type"))
state := c.GetString("state")
scope := c.GetString("scope")
if scope == "" {
scope = "openid"
}
challenge := c.GetString("code_challenge")
challengeMethod := c.GetString("code_challenge_method")
if challengeMethod == "" {
challengeMethod = "S256"
}
nonce := c.GetString("nonce")
if clientID == "" || redirectURI == "" {
c.Ctx.Output.SetStatus(400)
_, _ = c.Ctx.ResponseWriter.Write([]byte("缺少 client_id 或 redirect_uri"))
return
}
if responseType != ResponseTypeCode {
c.Ctx.Output.SetStatus(400)
_, _ = c.Ctx.ResponseWriter.Write([]byte("仅支持 response_type=code"))
return
}
client, err := findClient(clientID)
if err != nil {
c.Ctx.Output.SetStatus(400)
_, _ = c.Ctx.ResponseWriter.Write([]byte("client_id 无效"))
return
}
if !allowRedirect(client, redirectURI) {
c.Ctx.Output.SetStatus(400)
_, _ = c.Ctx.ResponseWriter.Write([]byte("redirect_uri 未登记"))
return
}
// 登录态:Cookie 优先(浏览器跳转),其次 Authorization(服务端调用)
sid := strings.TrimSpace(c.Ctx.GetCookie(sessionCookieName))
if sid == "" {
if claims := claimsFromHeader(c); claims != nil {
sid = claims.Sid
}
}
if sid == "" {
// 未登录 → 去登录页,登录后带着参数回来
back := fmt.Sprintf("%s?%s", authorizePath(c), c.Ctx.Request.URL.RawQuery)
target := fmt.Sprintf("/auth/login?redirect=%s&client_id=%s",
base64.RawURLEncoding.EncodeToString([]byte(back)), clientID)
c.Redirect(target, 302)
return
}
session, err := authsvc.GetSession(sid)
if err != nil {
clearSessionCookie(c)
back := fmt.Sprintf("%s?%s", authorizePath(c), c.Ctx.Request.URL.RawQuery)
target := fmt.Sprintf("/auth/login?redirect=%s&client_id=%s",
base64.RawURLEncoding.EncodeToString([]byte(back)), clientID)
c.Redirect(target, 302)
return
}
// 已登录但未选择企业:跳登录页的企业选择步骤
if session.Tid == authsvc.PendingTenantID {
back := fmt.Sprintf("%s?%s", authorizePath(c), c.Ctx.Request.URL.RawQuery)
target := fmt.Sprintf("/auth/login?step=tenant&redirect=%s&client_id=%s",
base64.RawURLEncoding.EncodeToString([]byte(back)), clientID)
c.Redirect(target, 302)
return
}
code, err := issueAuthCode(client.ClientID, session.IdentityID, session.Tid, redirectURI, challenge, challengeMethod, scope, nonce)
if err != nil {
c.Ctx.Output.SetStatus(500)
_, _ = c.Ctx.ResponseWriter.Write([]byte("签发授权码失败"))
return
}
sep := "?"
if strings.Contains(redirectURI, "?") {
sep = "&"
}
c.Redirect(fmt.Sprintf("%s%scode=%s&state=%s", redirectURI, sep, code, state), 302)
}
// Token 令牌端点
// POST /auth/token
// - grant_type=authorization_code:code + code_verifier(PKCE)+ client_id
// - grant_type=refresh_token:refresh_token + client_id
func (c *AuthOidcController) Token() {
grantType := strings.TrimSpace(c.GetString("grant_type"))
clientID := strings.TrimSpace(c.GetString("client_id"))
if clientID == "" {
// 兼容表单/JSON 以外的取参方式
clientID = strings.TrimSpace(c.Ctx.Request.FormValue("client_id"))
}
if grantType == "" {
grantType = strings.TrimSpace(c.Ctx.Request.FormValue("grant_type"))
}
switch grantType {
case GrantAuthCode:
code := strings.TrimSpace(c.GetString("code"))
verifier := strings.TrimSpace(c.GetString("code_verifier"))
if code == "" || verifier == "" || clientID == "" {
c.fail(400, "invalid_request")
return
}
pair, claims, err := exchangeCode(code, verifier, clientID)
if err != nil {
c.fail(400, err.Error())
return
}
idToken, _ := buildIDToken(claims, clientID)
c.serveJSON(map[string]interface{}{
"access_token": pair.AccessToken,
"refresh_token": pair.RefreshToken,
"token_type": pair.TokenType,
"expires_in": pair.ExpiresIn,
"id_token": idToken,
"sid": pair.Sid,
})
case GrantRefresh:
refresh := strings.TrimSpace(c.GetString("refresh_token"))
if refresh == "" || clientID == "" {
c.fail(400, "invalid_request")
return
}
pair, err := authsvc.RefreshTokens(refresh, clientID)
if err != nil {
c.fail(400, err.Error())
return
}
c.serveJSON(map[string]interface{}{
"access_token": pair.AccessToken,
"refresh_token": pair.RefreshToken,
"token_type": pair.TokenType,
"expires_in": pair.ExpiresIn,
"sid": pair.Sid,
})
default:
c.fail(400, "unsupported_grant_type")
}
}
// UserInfo 用户信息端点,需 Bearer Token
// GET /auth/userinfo
func (c *AuthOidcController) UserInfo() {
claims := claimsFromHeader(c)
if claims == nil {
c.Ctx.Output.SetStatus(401)
c.serveJSON(map[string]interface{}{"error": "invalid_token"})
return
}
if claims.UserID <= 0 {
c.fail(401, "invalid_token")
return
}
var identity models.AuthIdentity
if err := models.Orm.QueryTable(new(models.AuthIdentity)).
Filter("id", claims.UserID).One(&identity); err != nil {
c.fail(401, "invalid_token")
return
}
profile, err := authsvc.BuildProfile(&identity)
if err != nil {
c.fail(500, "server_error")
return
}
// 业务接口与前端统一使用 identity_id 作为用户标识
account, name, groupID := "", "", uint64(0)
if bind, err := authsvc.GetTenantUser(identity.ID, uint64(claims.TenantId)); err == nil {
groupID = bind.GroupID
if bind.Account != nil {
account = *bind.Account
}
if bind.Name != nil {
name = *bind.Name
}
}
if account == "" {
account = profile.Mobile
}
if name == "" {
name = profile.Nickname
}
// 当前会话所在企业名称(一人多企业时必须是"已选中的那一家",
// 而非全部可进入企业;前端直接展示该字段,避免自行拼接 tenants)
tenantName := ""
for _, t := range profile.Tenants {
if t.Tid == uint64(claims.TenantId) {
tenantName = t.TenantName
break
}
}
c.serveJSON(map[string]interface{}{
"sub": fmt.Sprintf("%d", identity.ID),
"id": identity.ID,
"union_id": identity.UnionID,
"tid": claims.TenantId,
"tenant_name": tenantName,
// 诊断用:业务接口按 user_type 判定权限(backend / app),
// 出现「无权访问」时可先看这里的值是否正确
"user_type": claims.UserType,
"group_id": groupID,
"account": account,
"name": name,
"nickname": profile.Nickname,
"mobile": profile.Mobile,
"email": profile.Email,
"avatar": profile.Avatar,
"tenants": profile.Tenants,
"client_id": claims.ClientID,
})
}
// Introspect 令牌校验,供资源服务(各业务后端)调用
// POST /auth/token 之外的独立端点;比本地 JWKS 验签更实时(可查黑名单)
func (c *AuthOidcController) Introspect() {
token := strings.TrimSpace(c.GetString("token"))
if token == "" {
token = strings.TrimSpace(c.Ctx.Request.FormValue("token"))
}
if token == "" {
c.fail(400, "invalid_request")
return
}
claims, err := jwtutil.ParseToken(token)
if err != nil {
c.serveJSON(map[string]interface{}{"active": false})
return
}
if authsvc.IsBlacklisted(claims.ID) {
c.serveJSON(map[string]interface{}{"active": false})
return
}
c.serveJSON(map[string]interface{}{
"active": true,
"sub": claims.Subject,
"user_id": claims.UserID,
"tid": claims.TenantId,
"client_id": claims.ClientID,
"sid": claims.Sid,
"scope": claims.Scope,
"amr": claims.Amr,
"exp": claims.ExpiresAt.Unix(),
})
}
// Revoke 吊销令牌(登出/踢下线)
// POST /auth/revoke
func (c *AuthOidcController) Revoke() {
token := strings.TrimSpace(c.GetString("token"))
refresh := strings.TrimSpace(c.GetString("refresh_token"))
if token == "" {
token = strings.TrimSpace(c.Ctx.Request.FormValue("token"))
}
if token == "" && refresh == "" {
c.fail(400, "invalid_request")
return
}
_ = authsvc.RevokeTokenPair(refresh, token, models.RevokeReasonLogout)
c.serveJSON(map[string]interface{}{"code": 200, "msg": "已吊销"})
}
// ---------------------------------------------------------------- 内部工具
func authorizePath(c *AuthOidcController) string {
return "/auth/authorize"
}
// clearSessionCookie 清除认证中心会话 Cookie
func clearSessionCookie(c *AuthOidcController) {
c.Ctx.Output.Header("Set-Cookie",
fmt.Sprintf("%s=; Path=/; Max-Age=0; HttpOnly; Secure; SameSite=Lax", sessionCookieName))
}
// setSessionCookie 写入认证中心会话 Cookie
func setSessionCookie(c *AuthOidcController, sid string) {
c.Ctx.Output.Header("Set-Cookie",
fmt.Sprintf("%s=%s; Path=/; Max-Age=%d; HttpOnly; Secure; SameSite=Lax",
sessionCookieName, sid, sessionCookieTTL))
}
func bearerToken(c *AuthOidcController) string {
header := c.Ctx.Request.Header.Get("Authorization")
if header == "" {
return ""
}
parts := strings.SplitN(header, " ", 2)
if len(parts) != 2 || parts[0] != "Bearer" {
return ""
}
return strings.TrimSpace(parts[1])
}
func claimsFromHeader(c *AuthOidcController) *jwtutil.Claims {
token := bearerToken(c)
if token == "" {
return nil
}
claims, err := jwtutil.ParseToken(token)
if err != nil {
return nil
}
return claims
}
// findClient 查询启用状态的应用
func findClient(clientID string) (*models.AuthClient, error) {
var client models.AuthClient
err := models.Orm.QueryTable(new(models.AuthClient)).
Filter("client_id", clientID).
Filter("status", 1).
One(&client)
if err != nil {
return nil, err
}
return &client, nil
}
// allowRedirect 校验回跳地址是否在白名单内(精确匹配,防钓鱼)
func allowRedirect(client *models.AuthClient, uri string) bool {
if client.RedirectURIs == nil || *client.RedirectURIs == "" {
return false
}
var list []string
if err := json.Unmarshal([]byte(*client.RedirectURIs), &list); err != nil {
return false
}
for _, item := range list {
if strings.TrimSpace(item) == strings.TrimSpace(uri) {
return true
}
}
return false
}
// allowLogoutRedirect 登出回跳地址校验。
//
// 先精确匹配白名单,再按 origin(协议+主机+端口)放宽匹配:
// 实际使用中「末尾斜杠」「带 #/login 片段」等差异很常见,
// 只做精确匹配会导致明明同域却跳不回去,因此同域即放行。
func allowLogoutRedirect(client *models.AuthClient, uri string) bool {
raw := ""
if client.PostLogoutURIs != nil {
raw = *client.PostLogoutURIs
}
var list []string
if raw != "" {
_ = json.Unmarshal([]byte(raw), &list)
}
if len(list) == 0 {
return false
}
target := strings.TrimSpace(uri)
for _, item := range list {
if strings.TrimSpace(item) == target {
return true
}
}
targetURL, err := url.Parse(target)
if err != nil || targetURL.Scheme == "" || targetURL.Host == "" {
return false
}
for _, item := range list {
base, err := url.Parse(strings.TrimSpace(item))
if err != nil || base.Scheme == "" || base.Host == "" {
continue
}
if base.Scheme == targetURL.Scheme && base.Host == targetURL.Host {
return true
}
}
return false
}
// issueAuthCode 生成一次性授权码(明文返回,库中只存哈希)
func issueAuthCode(clientID string, identityID, tid uint64, redirectURI, challenge, method, scope, nonce string) (string, error) {
plain, err := randomString(32)
if err != nil {
return "", err
}
sum := sha256.Sum256([]byte(plain))
code := &models.AuthCode{
CodeHash: fmt.Sprintf("%x", sum[:]),
ClientID: clientID,
IdentityID: identityID,
Tid: tid,
RedirectURI: redirectURI,
CodeChallenge: challenge,
CodeChallengeMethod: method,
ExpiresAt: time.Now().Add(60 * time.Second),
}
if scope != "" {
code.Scope = &scope
}
if nonce != "" {
code.Nonce = &nonce
}
if _, err := models.Orm.Insert(code); err != nil {
return "", err
}
return plain, nil
}
// exchangeCode 用授权码换令牌(校验 PKCE、一次性、有效期)
func exchangeCode(code, verifier, clientID string) (*authsvc.TokenPair, *jwtutil.Claims, error) {
sum := sha256.Sum256([]byte(code))
var stored models.AuthCode
if err := models.Orm.QueryTable(new(models.AuthCode)).
Filter("code_hash", fmt.Sprintf("%x", sum[:])).One(&stored); err != nil {
return nil, nil, fmt.Errorf("invalid_grant")
}
if stored.Used != 0 || stored.ExpiresAt.Before(time.Now()) {
return nil, nil, fmt.Errorf("invalid_grant")
}
if stored.ClientID != clientID {
return nil, nil, fmt.Errorf("invalid_client")
}
if !verifyPKCE(verifier, stored.CodeChallenge, stored.CodeChallengeMethod) {
return nil, nil, fmt.Errorf("invalid_grant")
}
// 一次性:立即标记已用
_, _ = models.Orm.QueryTable(new(models.AuthCode)).
Filter("code_hash", stored.CodeHash).
Update(map[string]interface{}{"used": 1})
session, err := authsvc.CreateSession(authsvc.SessionInfo{
IdentityID: stored.IdentityID,
Tid: stored.Tid,
ClientID: clientID,
LoginType: authsvc.LoginTypePassword,
Amr: authsvc.AmrPwd,
})
if err != nil {
return nil, nil, err
}
client, err := findClient(clientID)
accessTTL := 1800
refreshTTL := 2592000
if err == nil {
accessTTL = client.AccessTTL
refreshTTL = client.RefreshTTL
}
pair, err := authsvc.IssueTokens(authsvc.TokenIssue{
IdentityID: stored.IdentityID,
Tid: stored.Tid,
ClientID: clientID,
Sid: session.Sid,
UserType: "tenant",
Amr: authsvc.AmrPwd,
AccessTTL: accessTTL,
RefreshTTL: refreshTTL,
})
if err != nil {
return nil, nil, err
}
claims, err := jwtutil.ParseToken(pair.AccessToken)
if err != nil {
return nil, nil, err
}
return pair, claims, nil
}
// verifyPKCE 校验 PKCE(S256 或 plain)
func verifyPKCE(verifier, challenge, method string) bool {
if challenge == "" {
return false
}
if method == "S256" {
sum := sha256.Sum256([]byte(verifier))
return base64.RawURLEncoding.EncodeToString(sum[:]) == challenge
}
return verifier == challenge
}
// buildIDToken 生成 OIDC ID Token
func buildIDToken(claims *jwtutil.Claims, clientID string) (string, error) {
return jwtutil.SignToken(jwtutil.TokenOptions{
Alg: jwtutil.AlgRS256,
UserID: claims.UserID,
TenantID: claims.TenantId,
UserType: claims.UserType,
ClientID: clientID,
Sid: claims.Sid,
Subject: claims.Subject,
Audience: []string{clientID},
Amr: claims.Amr,
TTL: time.Hour,
})
}
// randomString 生成随机串
func randomString(n int) (string, error) {
buf := make([]byte, n)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(buf), nil
}