mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 14:06:36 +08:00
507 lines
15 KiB
Go
507 lines
15 KiB
Go
package service
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/rand"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/url"
|
|
"openflare/common"
|
|
"openflare/model"
|
|
"strings"
|
|
"time"
|
|
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
type PublicAuthSource struct {
|
|
ID uint `json:"id"`
|
|
Name string `json:"name"`
|
|
Type string `json:"type"`
|
|
DisplayName string `json:"display_name"`
|
|
AuthorizeURL string `json:"authorize_url"`
|
|
IconURL string `json:"icon_url"`
|
|
}
|
|
|
|
type OAuthProfile struct {
|
|
ExternalID string
|
|
ExternalUsername string
|
|
DisplayName string
|
|
Email string
|
|
}
|
|
|
|
type OAuthCallbackResult struct {
|
|
Status string `json:"status"`
|
|
User *model.User `json:"user,omitempty"`
|
|
}
|
|
|
|
type LinkExistingRequest struct {
|
|
Username string `json:"username"`
|
|
Password string `json:"password"`
|
|
}
|
|
|
|
type PendingExternalAccount struct {
|
|
AuthSourceID uint `json:"auth_source_id"`
|
|
ExternalID string `json:"external_id"`
|
|
ExternalUsername string `json:"external_username"`
|
|
DisplayName string `json:"display_name"`
|
|
Email string `json:"email"`
|
|
}
|
|
|
|
type oidcDiscovery struct {
|
|
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
|
TokenEndpoint string `json:"token_endpoint"`
|
|
UserInfoEndpoint string `json:"userinfo_endpoint"`
|
|
JWKSURI string `json:"jwks_uri"`
|
|
Issuer string `json:"issuer"`
|
|
}
|
|
|
|
type oauthTokenResponse struct {
|
|
AccessToken string `json:"access_token"`
|
|
TokenType string `json:"token_type"`
|
|
IDToken string `json:"id_token"`
|
|
Scope string `json:"scope"`
|
|
}
|
|
|
|
var oauthHTTPClient = &http.Client{Timeout: 8 * time.Second}
|
|
|
|
func GenerateOAuthState() (string, error) {
|
|
buffer := make([]byte, 24)
|
|
if _, err := rand.Read(buffer); err != nil {
|
|
return "", err
|
|
}
|
|
return base64.RawURLEncoding.EncodeToString(buffer), nil
|
|
}
|
|
|
|
func PublicAuthSources(baseAPIPath string) ([]PublicAuthSource, error) {
|
|
sources, err := model.GetActiveAuthSources()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
result := make([]PublicAuthSource, 0, len(sources))
|
|
for _, source := range sources {
|
|
result = append(result, PublicAuthSource{
|
|
ID: source.ID,
|
|
Name: source.Name,
|
|
Type: source.Type,
|
|
DisplayName: source.DisplayName,
|
|
AuthorizeURL: fmt.Sprintf("%s/oauth/%s/authorize", strings.TrimRight(baseAPIPath, "/"), url.PathEscape(source.Name)),
|
|
IconURL: source.IconURL,
|
|
})
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
func BuildAuthorizeURL(ctx context.Context, source *model.AuthSource, redirectURL string, state string) (string, error) {
|
|
source.Normalize()
|
|
switch source.Type {
|
|
case model.AuthSourceTypeGitHub:
|
|
authorizeURL, err := url.Parse("https://github.com/login/oauth/authorize")
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
values := authorizeURL.Query()
|
|
values.Set("client_id", source.ClientID)
|
|
values.Set("redirect_uri", redirectURL)
|
|
values.Set("scope", source.Scopes)
|
|
values.Set("state", state)
|
|
authorizeURL.RawQuery = values.Encode()
|
|
return authorizeURL.String(), nil
|
|
case model.AuthSourceTypeOIDC:
|
|
discovery, err := fetchOIDCDiscovery(ctx, source.OpenIDDiscoveryURL)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
authorizeURL, err := url.Parse(discovery.AuthorizationEndpoint)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
values := authorizeURL.Query()
|
|
values.Set("client_id", source.ClientID)
|
|
values.Set("redirect_uri", redirectURL)
|
|
values.Set("response_type", "code")
|
|
values.Set("scope", source.Scopes)
|
|
values.Set("state", state)
|
|
authorizeURL.RawQuery = values.Encode()
|
|
return authorizeURL.String(), nil
|
|
default:
|
|
return "", errors.New("不支持的认证源类型")
|
|
}
|
|
}
|
|
|
|
func ExchangeOAuthProfile(ctx context.Context, source *model.AuthSource, code string, redirectURL string) (*OAuthProfile, error) {
|
|
if strings.TrimSpace(code) == "" {
|
|
return nil, errors.New("授权 code 不能为空")
|
|
}
|
|
source.Normalize()
|
|
switch source.Type {
|
|
case model.AuthSourceTypeGitHub:
|
|
return exchangeGitHubProfile(ctx, source, code, redirectURL)
|
|
case model.AuthSourceTypeOIDC:
|
|
return exchangeOIDCProfile(ctx, source, code, redirectURL)
|
|
default:
|
|
return nil, errors.New("不支持的认证源类型")
|
|
}
|
|
}
|
|
|
|
func CompleteOAuthLogin(source *model.AuthSource, profile *OAuthProfile, currentUserID *int) (*OAuthCallbackResult, *PendingExternalAccount, error) {
|
|
if source == nil || profile == nil || strings.TrimSpace(profile.ExternalID) == "" {
|
|
return nil, nil, errors.New("第三方账号资料不完整")
|
|
}
|
|
|
|
account, err := model.FindExternalAccount(source.ID, profile.ExternalID)
|
|
if err == nil {
|
|
user, err := model.GetUserById(account.UserID, false)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
if user.Status != common.UserStatusEnabled {
|
|
return nil, nil, errors.New("用户已被封禁")
|
|
}
|
|
return &OAuthCallbackResult{Status: "logged_in", User: user}, nil, nil
|
|
}
|
|
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return nil, nil, err
|
|
}
|
|
|
|
if currentUserID != nil && *currentUserID > 0 {
|
|
user, err := model.GetUserById(*currentUserID, false)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
if user.Status != common.UserStatusEnabled {
|
|
return nil, nil, errors.New("用户已被封禁")
|
|
}
|
|
if err := model.LinkExternalAccount(&model.ExternalAccount{
|
|
AuthSourceID: source.ID,
|
|
UserID: user.Id,
|
|
ExternalID: profile.ExternalID,
|
|
ExternalUsername: profile.ExternalUsername,
|
|
Email: profile.Email,
|
|
}); err != nil {
|
|
return nil, nil, err
|
|
}
|
|
return &OAuthCallbackResult{Status: "linked", User: user}, nil, nil
|
|
}
|
|
|
|
if common.RegisterEnabled {
|
|
user, err := createUserFromOAuthProfile(source, profile)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
return &OAuthCallbackResult{Status: "registered", User: user}, nil, nil
|
|
}
|
|
|
|
pending := &PendingExternalAccount{
|
|
AuthSourceID: source.ID,
|
|
ExternalID: profile.ExternalID,
|
|
ExternalUsername: profile.ExternalUsername,
|
|
DisplayName: profile.DisplayName,
|
|
Email: profile.Email,
|
|
}
|
|
return &OAuthCallbackResult{Status: "link_required"}, pending, nil
|
|
}
|
|
|
|
func LinkPendingExternalAccount(pending *PendingExternalAccount, input LinkExistingRequest) (*model.User, error) {
|
|
if pending == nil || pending.AuthSourceID == 0 || pending.ExternalID == "" {
|
|
return nil, errors.New("待绑定第三方账号已失效,请重新登录")
|
|
}
|
|
user := model.User{
|
|
Username: strings.TrimSpace(input.Username),
|
|
Password: input.Password,
|
|
}
|
|
if err := user.ValidateAndFill(); err != nil {
|
|
return nil, err
|
|
}
|
|
if user.Status != common.UserStatusEnabled {
|
|
return nil, errors.New("用户已被封禁")
|
|
}
|
|
|
|
if existing, err := model.FindExternalAccount(pending.AuthSourceID, pending.ExternalID); err == nil {
|
|
if existing.UserID != user.Id {
|
|
return nil, errors.New("该第三方账号已绑定其他用户")
|
|
}
|
|
return &user, nil
|
|
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return nil, err
|
|
}
|
|
|
|
if err := model.LinkExternalAccount(&model.ExternalAccount{
|
|
AuthSourceID: pending.AuthSourceID,
|
|
UserID: user.Id,
|
|
ExternalID: pending.ExternalID,
|
|
ExternalUsername: pending.ExternalUsername,
|
|
Email: pending.Email,
|
|
}); err != nil {
|
|
return nil, err
|
|
}
|
|
return &user, nil
|
|
}
|
|
|
|
func createUserFromOAuthProfile(source *model.AuthSource, profile *OAuthProfile) (*model.User, error) {
|
|
displayName := strings.TrimSpace(profile.DisplayName)
|
|
if displayName == "" {
|
|
displayName = strings.TrimSpace(profile.ExternalUsername)
|
|
}
|
|
if displayName == "" {
|
|
displayName = source.DisplayName + " User"
|
|
}
|
|
if len([]rune(displayName)) > 20 {
|
|
displayName = string([]rune(displayName)[:20])
|
|
}
|
|
|
|
prefix := source.Type
|
|
if prefix == "" {
|
|
prefix = "oauth"
|
|
}
|
|
var username string
|
|
for index := 0; index < 20; index++ {
|
|
username = fmt.Sprintf("%s_%d", prefix, model.GetMaxUserId()+1+index)
|
|
if !model.IsUsernameAlreadyTaken(username) {
|
|
break
|
|
}
|
|
}
|
|
|
|
user := &model.User{
|
|
Username: username,
|
|
DisplayName: displayName,
|
|
Email: profile.Email,
|
|
Role: common.RoleCommonUser,
|
|
Status: common.UserStatusEnabled,
|
|
}
|
|
if err := user.Insert(); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := model.LinkExternalAccount(&model.ExternalAccount{
|
|
AuthSourceID: source.ID,
|
|
UserID: user.Id,
|
|
ExternalID: profile.ExternalID,
|
|
ExternalUsername: profile.ExternalUsername,
|
|
Email: profile.Email,
|
|
}); err != nil {
|
|
return nil, err
|
|
}
|
|
return user, nil
|
|
}
|
|
|
|
func exchangeGitHubProfile(ctx context.Context, source *model.AuthSource, code string, redirectURL string) (*OAuthProfile, error) {
|
|
values := map[string]string{
|
|
"client_id": source.ClientID,
|
|
"client_secret": source.ClientSecret,
|
|
"code": code,
|
|
"redirect_uri": redirectURL,
|
|
}
|
|
body, err := json.Marshal(values)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, "https://github.com/login/oauth/access_token", bytes.NewReader(body))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
req.Header.Set("Content-Type", "application/json")
|
|
req.Header.Set("Accept", "application/json")
|
|
resp, err := oauthHTTPClient.Do(req)
|
|
if err != nil {
|
|
slog.Error("github oauth access token request failed", "error", err)
|
|
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试")
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
return nil, fmt.Errorf("GitHub token 接口返回异常状态: %s", resp.Status)
|
|
}
|
|
var token oauthTokenResponse
|
|
if err := json.NewDecoder(resp.Body).Decode(&token); err != nil {
|
|
return nil, err
|
|
}
|
|
if token.AccessToken == "" {
|
|
return nil, errors.New("GitHub 未返回 access token")
|
|
}
|
|
|
|
req, err = http.NewRequestWithContext(ctx, http.MethodGet, "https://api.github.com/user", nil)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+token.AccessToken)
|
|
req.Header.Set("Accept", "application/vnd.github+json")
|
|
resp, err = oauthHTTPClient.Do(req)
|
|
if err != nil {
|
|
slog.Error("github user info request failed", "error", err)
|
|
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试")
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
return nil, fmt.Errorf("GitHub 用户接口返回异常状态: %s", resp.Status)
|
|
}
|
|
var githubUser struct {
|
|
ID int64 `json:"id"`
|
|
Login string `json:"login"`
|
|
Name string `json:"name"`
|
|
Email string `json:"email"`
|
|
}
|
|
if err := json.NewDecoder(resp.Body).Decode(&githubUser); err != nil {
|
|
return nil, err
|
|
}
|
|
if githubUser.ID == 0 && githubUser.Login == "" {
|
|
return nil, errors.New("GitHub 用户资料缺少唯一标识")
|
|
}
|
|
return &OAuthProfile{
|
|
ExternalID: githubUser.Login,
|
|
ExternalUsername: githubUser.Login,
|
|
DisplayName: firstNonEmpty(githubUser.Name, githubUser.Login),
|
|
Email: githubUser.Email,
|
|
}, nil
|
|
}
|
|
|
|
func exchangeOIDCProfile(ctx context.Context, source *model.AuthSource, code string, redirectURL string) (*OAuthProfile, error) {
|
|
discovery, err := fetchOIDCDiscovery(ctx, source.OpenIDDiscoveryURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
token, err := exchangeOIDCToken(ctx, discovery.TokenEndpoint, source, code, redirectURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if token.AccessToken == "" {
|
|
return nil, errors.New("OIDC 未返回 access token")
|
|
}
|
|
claims, err := fetchOIDCUserInfo(ctx, discovery.UserInfoEndpoint, token.AccessToken)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(claims) == 0 && token.IDToken != "" {
|
|
claims = decodeJWTClaims(token.IDToken)
|
|
}
|
|
profile := profileFromClaims(claims)
|
|
if profile.ExternalID == "" {
|
|
return nil, errors.New("OIDC 用户资料缺少 sub")
|
|
}
|
|
return profile, nil
|
|
}
|
|
|
|
func fetchOIDCDiscovery(ctx context.Context, discoveryURL string) (*oidcDiscovery, error) {
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, discoveryURL, nil)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
resp, err := oauthHTTPClient.Do(req)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("无法获取 OIDC discovery 配置: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
return nil, fmt.Errorf("OIDC discovery 返回异常状态: %s", resp.Status)
|
|
}
|
|
var discovery oidcDiscovery
|
|
if err := json.NewDecoder(resp.Body).Decode(&discovery); err != nil {
|
|
return nil, err
|
|
}
|
|
if discovery.AuthorizationEndpoint == "" || discovery.TokenEndpoint == "" {
|
|
return nil, errors.New("OIDC discovery 缺少授权或 token 端点")
|
|
}
|
|
return &discovery, nil
|
|
}
|
|
|
|
func exchangeOIDCToken(ctx context.Context, tokenEndpoint string, source *model.AuthSource, code string, redirectURL string) (*oauthTokenResponse, error) {
|
|
form := url.Values{}
|
|
form.Set("grant_type", "authorization_code")
|
|
form.Set("client_id", source.ClientID)
|
|
form.Set("client_secret", source.ClientSecret)
|
|
form.Set("code", code)
|
|
form.Set("redirect_uri", redirectURL)
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, tokenEndpoint, strings.NewReader(form.Encode()))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
req.Header.Set("Accept", "application/json")
|
|
resp, err := oauthHTTPClient.Do(req)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("OIDC token 请求失败: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
|
return nil, fmt.Errorf("OIDC token 接口返回异常状态: %s %s", resp.Status, strings.TrimSpace(string(raw)))
|
|
}
|
|
var token oauthTokenResponse
|
|
if err := json.NewDecoder(resp.Body).Decode(&token); err != nil {
|
|
return nil, err
|
|
}
|
|
return &token, nil
|
|
}
|
|
|
|
func fetchOIDCUserInfo(ctx context.Context, endpoint string, accessToken string) (map[string]any, error) {
|
|
if endpoint == "" {
|
|
return map[string]any{}, nil
|
|
}
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+accessToken)
|
|
req.Header.Set("Accept", "application/json")
|
|
resp, err := oauthHTTPClient.Do(req)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("OIDC userinfo 请求失败: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
|
return nil, fmt.Errorf("OIDC userinfo 返回异常状态: %s %s", resp.Status, strings.TrimSpace(string(raw)))
|
|
}
|
|
var claims map[string]any
|
|
if err := json.NewDecoder(resp.Body).Decode(&claims); err != nil {
|
|
return nil, err
|
|
}
|
|
return claims, nil
|
|
}
|
|
|
|
func decodeJWTClaims(token string) map[string]any {
|
|
parts := strings.Split(token, ".")
|
|
if len(parts) < 2 {
|
|
return map[string]any{}
|
|
}
|
|
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
|
|
if err != nil {
|
|
return map[string]any{}
|
|
}
|
|
var claims map[string]any
|
|
if err := json.Unmarshal(payload, &claims); err != nil {
|
|
return map[string]any{}
|
|
}
|
|
return claims
|
|
}
|
|
|
|
func profileFromClaims(claims map[string]any) *OAuthProfile {
|
|
stringClaim := func(keys ...string) string {
|
|
for _, key := range keys {
|
|
if value, ok := claims[key].(string); ok && strings.TrimSpace(value) != "" {
|
|
return strings.TrimSpace(value)
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
return &OAuthProfile{
|
|
ExternalID: stringClaim("sub"),
|
|
ExternalUsername: stringClaim("preferred_username", "nickname", "name", "email"),
|
|
DisplayName: stringClaim("name", "preferred_username", "nickname", "email"),
|
|
Email: stringClaim("email"),
|
|
}
|
|
}
|
|
|
|
func firstNonEmpty(values ...string) string {
|
|
for _, value := range values {
|
|
if strings.TrimSpace(value) != "" {
|
|
return strings.TrimSpace(value)
|
|
}
|
|
}
|
|
return ""
|
|
}
|