// 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" "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 } c.serveJSON(map[string]interface{}{ "sub": fmt.Sprintf("%d", identity.ID), "union_id": identity.UnionID, "tid": claims.TenantId, "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 } // 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 }