mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 16:46:37 +08:00
去除默认OIDC
This commit is contained in:
@@ -1,73 +0,0 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"strings"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
var (
|
||||
oauthConf *oauth2.Config
|
||||
oidcVerifier *oidc.IDTokenVerifier
|
||||
)
|
||||
|
||||
func init() {
|
||||
cfg := config.Config.OAuth2
|
||||
|
||||
if cfg.Issuer != "" {
|
||||
ctx := context.Background()
|
||||
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
|
||||
issuer := strings.TrimSuffix(strings.TrimSpace(cfg.Issuer), "/")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
|
||||
|
||||
provider, err := oidc.NewProvider(ctx, issuer)
|
||||
if err != nil {
|
||||
log.Printf("[OAuth] 初始化 OIDC Provider 失败: %v,将仅使用 OAuth2", err)
|
||||
} else {
|
||||
oidcVerifier = provider.Verifier(&oidc.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
})
|
||||
log.Printf("[OAuth] OIDC Provider 初始化成功: %s", issuer)
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化 OAuth2 配置
|
||||
scopes := []string{"profile", "email"}
|
||||
if oidcVerifier != nil {
|
||||
// 启用 OIDC 时添加 openid scope
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
|
||||
oauthConf = &oauth2.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
ClientSecret: cfg.ClientSecret,
|
||||
RedirectURL: cfg.RedirectURI,
|
||||
Scopes: scopes,
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: cfg.AuthorizationEndpoint,
|
||||
TokenURL: cfg.TokenEndpoint,
|
||||
AuthStyle: oauth2.AuthStyleAutoDetect,
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -1,152 +0,0 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
"github.com/linux-do/credit/internal/common"
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/otel_trace"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func GetUserIDFromSession(s sessions.Session) uint64 {
|
||||
userID, ok := s.Get(UserIDKey).(uint64)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return userID
|
||||
}
|
||||
|
||||
func GetUserIDFromContext(c *gin.Context) uint64 {
|
||||
session := sessions.Default(c)
|
||||
return GetUserIDFromSession(session)
|
||||
}
|
||||
|
||||
// doOAuth 执行 OAuth2/OIDC 流程
|
||||
func doOAuth(ctx context.Context, code string, nonce string) (*model.User, error) {
|
||||
ctx, span := otel_trace.Start(ctx, "OAuth")
|
||||
defer span.End()
|
||||
|
||||
var userInfo model.OAuthUserInfo
|
||||
|
||||
// 使用授权码换取 Token
|
||||
token, err := oauthConf.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if oidcVerifier != nil {
|
||||
if rawIDToken, ok := token.Extra("id_token").(string); ok {
|
||||
idToken, verifyErr := oidcVerifier.Verify(ctx, rawIDToken)
|
||||
if verifyErr != nil {
|
||||
err := fmt.Errorf("%s: %w", IDTokenVerifyFailed, verifyErr)
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
if nonce != "" && idToken.Nonce != nonce {
|
||||
span.SetStatus(codes.Error, NonceMismatch)
|
||||
return nil, errors.New(NonceMismatch)
|
||||
}
|
||||
if claimsErr := idToken.Claims(&userInfo); claimsErr != nil {
|
||||
span.SetStatus(codes.Error, claimsErr.Error())
|
||||
return nil, claimsErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if userInfo.GetID() == 0 {
|
||||
client := oauthConf.Client(ctx, token)
|
||||
resp, httpErr := client.Get(config.Config.OAuth2.UserEndpoint)
|
||||
if httpErr != nil {
|
||||
span.SetStatus(codes.Error, httpErr.Error())
|
||||
return nil, httpErr
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
responseData, readErr := io.ReadAll(resp.Body)
|
||||
if readErr != nil {
|
||||
span.SetStatus(codes.Error, readErr.Error())
|
||||
return nil, readErr
|
||||
}
|
||||
if unmarshalErr := json.Unmarshal(responseData, &userInfo); unmarshalErr != nil {
|
||||
span.SetStatus(codes.Error, unmarshalErr.Error())
|
||||
return nil, unmarshalErr
|
||||
}
|
||||
}
|
||||
|
||||
if !userInfo.Active {
|
||||
err := errors.New(common.BannedAccount)
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var user model.User
|
||||
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var holder model.User
|
||||
if conflictErr := tx.Where("username = ? AND id != ?", userInfo.Username, userInfo.GetID()).First(&holder).Error; conflictErr == nil {
|
||||
// 存在冲突 -> 将占用者改名并注销
|
||||
newParams := map[string]interface{}{
|
||||
"username": fmt.Sprintf("%s已注销: %s", holder.Username, uuid.NewString()),
|
||||
"is_active": false,
|
||||
}
|
||||
if updateErr := tx.Model(&holder).Updates(newParams).Error; updateErr != nil {
|
||||
return updateErr
|
||||
}
|
||||
}
|
||||
|
||||
// 根据 ID 处理当前用户的 更新 或 创建
|
||||
if queryErr := tx.Where("id = ?", userInfo.GetID()).First(&user).Error; queryErr == nil {
|
||||
// 用户已存在 -> 更新信息
|
||||
if activeErr := user.CheckActive(); activeErr != nil {
|
||||
return activeErr
|
||||
}
|
||||
user.UpdateFromOAuthInfo(&userInfo)
|
||||
if saveErr := tx.Save(&user).Error; saveErr != nil {
|
||||
return saveErr
|
||||
}
|
||||
} else if errors.Is(queryErr, gorm.ErrRecordNotFound) {
|
||||
// 用户不存在 -> 创建新用户
|
||||
user = model.User{}
|
||||
if createErr := user.CreateUser(tx, &userInfo); createErr != nil {
|
||||
return createErr
|
||||
}
|
||||
} else {
|
||||
return queryErr
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
}
|
||||
@@ -27,6 +27,7 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -134,6 +135,17 @@ var (
|
||||
testJWKS jose.JSONWebKeySet
|
||||
)
|
||||
|
||||
const (
|
||||
testIssuerURL = "https://connect.linux.do"
|
||||
testAuthURL = "https://connect.linux.do/oauth2/authorize"
|
||||
testTokenURL = "https://connect.linux.do/oauth2/token"
|
||||
testJWKSURL = "https://connect.linux.do/oauth2/keys"
|
||||
testClientID = "test_client_id"
|
||||
testClientSecret = "test_client_secret"
|
||||
testSourceName = "linuxdo"
|
||||
testSourceDisplay = "LINUX DO"
|
||||
)
|
||||
|
||||
func init() {
|
||||
var err error
|
||||
testRSAPrivateKey, err = rsa.GenerateKey(rand.Reader, 2048)
|
||||
@@ -151,7 +163,50 @@ func init() {
|
||||
}
|
||||
}
|
||||
|
||||
func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) {
|
||||
t.Helper()
|
||||
if err := dbConn.Create(&model.AuthSource{
|
||||
ID: 100,
|
||||
Name: testSourceName,
|
||||
Type: model.AuthSourceTypeOIDC,
|
||||
DisplayName: testSourceDisplay,
|
||||
IsActive: true,
|
||||
ClientID: testClientID,
|
||||
ClientSecret: testClientSecret,
|
||||
OpenIDDiscoveryURL: testIssuerURL,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed auth source: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func oidcDiscoveryResponse() *http.Response {
|
||||
body := fmt.Sprintf(`{
|
||||
"issuer": %q,
|
||||
"authorization_endpoint": %q,
|
||||
"token_endpoint": %q,
|
||||
"jwks_uri": %q,
|
||||
"response_types_supported": ["code"],
|
||||
"subject_types_supported": ["public"],
|
||||
"id_token_signing_alg_values_supported": ["RS256"]
|
||||
}`, testIssuerURL, testAuthURL, testTokenURL, testJWKSURL)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
Header: make(http.Header),
|
||||
}
|
||||
}
|
||||
|
||||
func jwksResponse() *http.Response {
|
||||
jwksJSON, _ := json.Marshal(testJWKS)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
|
||||
Header: make(http.Header),
|
||||
}
|
||||
}
|
||||
|
||||
type mockClaims struct {
|
||||
ID uint64 `json:"id"`
|
||||
Issuer string `json:"iss"`
|
||||
Subject string `json:"sub"`
|
||||
Audience string `json:"aud"`
|
||||
@@ -169,7 +224,9 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
id, _ := strconv.ParseUint(sub, 10, 64)
|
||||
claims := mockClaims{
|
||||
ID: id,
|
||||
Issuer: issuer,
|
||||
Subject: sub,
|
||||
Audience: aud,
|
||||
@@ -236,19 +293,6 @@ func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *ht
|
||||
db.SetDB(dbConn)
|
||||
db.Redis = mockRedis
|
||||
|
||||
// Setup oauth vars
|
||||
oauthConf = &oauth2.Config{
|
||||
ClientID: config.Config.OAuth2.ClientID,
|
||||
ClientSecret: config.Config.OAuth2.ClientSecret,
|
||||
RedirectURL: config.Config.OAuth2.RedirectURI,
|
||||
Scopes: []string{"profile", "email"},
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: config.Config.OAuth2.AuthorizationEndpoint,
|
||||
TokenURL: config.Config.OAuth2.TokenEndpoint,
|
||||
},
|
||||
}
|
||||
oidcVerifier = nil
|
||||
|
||||
api := r.Group("/api/v1")
|
||||
{
|
||||
api.GET("/oauth/sources", GetLoginSources)
|
||||
@@ -288,12 +332,7 @@ func initializeTestConfig() {
|
||||
config.Config.App.SessionCookieName = "test_session_id"
|
||||
config.Config.App.SessionSecret = "test_session_secret"
|
||||
config.Config.App.APIPrefix = "/api"
|
||||
config.Config.OAuth2.ClientID = "test_client_id"
|
||||
config.Config.OAuth2.ClientSecret = "test_client_secret"
|
||||
config.Config.OAuth2.RedirectURI = "http://localhost:3000/callback"
|
||||
config.Config.OAuth2.AuthorizationEndpoint = "https://connect.linux.do/oauth2/authorize"
|
||||
config.Config.OAuth2.TokenEndpoint = "https://connect.linux.do/oauth2/token"
|
||||
config.Config.OAuth2.UserEndpoint = "https://connect.linux.do/api/user"
|
||||
config.Config.App.FrontendURL = "http://localhost:3000"
|
||||
config.Config.OpenAPIRisk.Enabled = false
|
||||
config.Config.OpenAPIRisk.BaseURL = ""
|
||||
config.Config.OpenAPIRisk.BlockRiskLevels = []string{}
|
||||
@@ -352,13 +391,12 @@ func TestGetLoginSources(t *testing.T) {
|
||||
t.Fatalf("failed to unmarshal response: %v", err)
|
||||
}
|
||||
|
||||
if len(resp.Data) != 2 {
|
||||
t.Fatalf("expected 2 active sources, got %d", len(resp.Data))
|
||||
if len(resp.Data) != 1 {
|
||||
t.Fatalf("expected 1 active source, got %d", len(resp.Data))
|
||||
}
|
||||
|
||||
names := []string{resp.Data[0].Name, resp.Data[1].Name}
|
||||
if !strings.Contains(strings.Join(names, ","), "default") || !strings.Contains(strings.Join(names, ","), "github") {
|
||||
t.Errorf("expected default and github sources, got %v", names)
|
||||
if resp.Data[0].Name != "github" {
|
||||
t.Errorf("expected github source, got %s", resp.Data[0].Name)
|
||||
}
|
||||
|
||||
// Test disabling OIDC
|
||||
@@ -379,10 +417,14 @@ func TestGetLoginURL(t *testing.T) {
|
||||
initializeTestConfig()
|
||||
dbConn := setupTestDB(t)
|
||||
mockRedis := newMockRedisClient()
|
||||
seedTestAuthSource(t, dbConn)
|
||||
|
||||
httpMock := &http.Client{
|
||||
Transport: &mockRoundTripper{
|
||||
roundTripFunc: func(req *http.Request) (*http.Response, error) {
|
||||
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
|
||||
return oidcDiscoveryResponse(), nil
|
||||
}
|
||||
return nil, fmt.Errorf("unexpected request")
|
||||
},
|
||||
},
|
||||
@@ -404,7 +446,7 @@ func TestGetLoginURL(t *testing.T) {
|
||||
t.Fatalf("failed to unmarshal response: %v", err)
|
||||
}
|
||||
|
||||
if !strings.Contains(resp.Data.AuthorizeURL, config.Config.OAuth2.AuthorizationEndpoint) {
|
||||
if !strings.Contains(resp.Data.AuthorizeURL, testAuthURL) {
|
||||
t.Errorf("invalid authorize URL: %s", resp.Data.AuthorizeURL)
|
||||
}
|
||||
|
||||
@@ -424,7 +466,7 @@ func TestGetLoginURL(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("failed to decode state payload: %v", err)
|
||||
}
|
||||
if payload.SourceName != "default" || payload.Purpose != OAuthPurposeLogin {
|
||||
if payload.SourceName != testSourceName || payload.Purpose != OAuthPurposeLogin {
|
||||
t.Errorf("unexpected payload: %+v", payload)
|
||||
}
|
||||
|
||||
@@ -525,23 +567,23 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
initializeTestConfig()
|
||||
dbConn := setupTestDB(t)
|
||||
mockRedis := newMockRedisClient()
|
||||
seedTestAuthSource(t, dbConn)
|
||||
state := uuid.NewString()
|
||||
|
||||
// 1. Mock the outgoing HTTP client for token exchange and user info fetching
|
||||
httpMock := &http.Client{
|
||||
Transport: &mockRoundTripper{
|
||||
roundTripFunc: func(req *http.Request) (*http.Response, error) {
|
||||
// Handle Token Exchange
|
||||
if req.Method == http.MethodPost && req.URL.String() == config.Config.OAuth2.TokenEndpoint {
|
||||
body := `{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600}`
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
Header: make(http.Header),
|
||||
}, nil
|
||||
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
|
||||
return oidcDiscoveryResponse(), nil
|
||||
}
|
||||
// Handle User Info fetch
|
||||
if req.Method == http.MethodGet && req.URL.String() == config.Config.OAuth2.UserEndpoint {
|
||||
body := `{"id":88888,"sub":"oauth_user_88888","username":"test_oauth_user","email":"oauth@linux.do","name":"Oauth Test User","active":true,"trust_level":3}`
|
||||
if req.Method == http.MethodGet && req.URL.String() == testJWKSURL {
|
||||
return jwksResponse(), nil
|
||||
}
|
||||
// Handle Token Exchange
|
||||
if req.Method == http.MethodPost && req.URL.String() == testTokenURL {
|
||||
idToken := generateMockIDToken(testIssuerURL, "88888", testClientID, state, "test_oauth_user", "oauth@linux.do", "Oauth Test User")
|
||||
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
@@ -556,9 +598,8 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
router := setupTestRouter(dbConn, mockRedis, httpMock)
|
||||
|
||||
// 2. Setup state in Redis
|
||||
state := uuid.NewString()
|
||||
payloadValue, _ := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: "default",
|
||||
SourceName: testSourceName,
|
||||
Purpose: OAuthPurposeLogin,
|
||||
})
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration)
|
||||
@@ -582,13 +623,13 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
t.Errorf("expected logged_in status, got %s", callbackResp.Data.Status)
|
||||
}
|
||||
|
||||
if callbackResp.Data.User.Username != "test_oauth_user" || callbackResp.Data.User.ID != 1 {
|
||||
if callbackResp.Data.User.Username != "test_oauth_user" || callbackResp.Data.User.ID != 88888 {
|
||||
t.Errorf("unexpected user returned: %+v", callbackResp.Data.User)
|
||||
}
|
||||
|
||||
// Verify user is created in database
|
||||
var user model.User
|
||||
if err := dbConn.First(&user, "id = ?", 1).Error; err != nil {
|
||||
if err := dbConn.First(&user, "id = ?", 88888).Error; err != nil {
|
||||
t.Fatalf("user was not created in DB: %v", err)
|
||||
}
|
||||
|
||||
@@ -619,15 +660,15 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
httpMock2 := &http.Client{
|
||||
Transport: &mockRoundTripper{
|
||||
roundTripFunc: func(req *http.Request) (*http.Response, error) {
|
||||
if req.Method == http.MethodPost && req.URL.String() == config.Config.OAuth2.TokenEndpoint {
|
||||
body := `{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600}`
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
}, nil
|
||||
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
|
||||
return oidcDiscoveryResponse(), nil
|
||||
}
|
||||
if req.Method == http.MethodGet && req.URL.String() == config.Config.OAuth2.UserEndpoint {
|
||||
body := `{"id":99999,"sub":"oauth_user_99999","username":"test_oauth_user","email":"another@linux.do","name":"Another User","active":true,"trust_level":1}`
|
||||
if req.Method == http.MethodGet && req.URL.String() == testJWKSURL {
|
||||
return jwksResponse(), nil
|
||||
}
|
||||
if req.Method == http.MethodPost && req.URL.String() == testTokenURL {
|
||||
idToken := generateMockIDToken(testIssuerURL, "99999", testClientID, state2, "test_oauth_user", "another@linux.do", "Another User")
|
||||
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
|
||||
@@ -2,10 +2,8 @@ package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -53,33 +51,30 @@ type CallbackRequest struct {
|
||||
Code string `json:"code" binding:"required"`
|
||||
}
|
||||
|
||||
func defaultAuthSource() *model.AuthSource {
|
||||
if config.Config.OAuth2.ClientID == "" || config.Config.OAuth2.RedirectURI == "" {
|
||||
return nil
|
||||
func GetUserIDFromSession(s sessions.Session) uint64 {
|
||||
userID, ok := s.Get(UserIDKey).(uint64)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
source := &model.AuthSource{
|
||||
Name: "default",
|
||||
Type: model.AuthSourceTypeOIDC,
|
||||
DisplayName: "默认认证源",
|
||||
IsActive: true,
|
||||
ClientID: config.Config.OAuth2.ClientID,
|
||||
ClientSecret: config.Config.OAuth2.ClientSecret,
|
||||
OpenIDDiscoveryURL: config.Config.OAuth2.Issuer,
|
||||
}
|
||||
if source.DisplayName == "" {
|
||||
source.DisplayName = "默认认证源"
|
||||
}
|
||||
return source
|
||||
return userID
|
||||
}
|
||||
|
||||
func GetUserIDFromContext(c *gin.Context) uint64 {
|
||||
session := sessions.Default(c)
|
||||
return GetUserIDFromSession(session)
|
||||
}
|
||||
|
||||
func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
|
||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||
if name == "" || name == "default" {
|
||||
source := defaultAuthSource()
|
||||
if source == nil {
|
||||
return nil, errors.New("默认认证源未配置")
|
||||
if name == "" {
|
||||
sources, err := model.GetActiveAuthSources()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return source, nil
|
||||
if len(sources) == 0 {
|
||||
return nil, errors.New("未配置可用认证源")
|
||||
}
|
||||
return &sources[0], nil
|
||||
}
|
||||
return model.GetAuthSourceByName(name)
|
||||
}
|
||||
@@ -90,24 +85,11 @@ func activeLoginSources() []AuthSourceView {
|
||||
return nil
|
||||
}
|
||||
|
||||
sources := make([]AuthSourceView, 0, 4)
|
||||
if source := defaultAuthSource(); source != nil {
|
||||
source.Sanitize()
|
||||
sources = append(sources, AuthSourceView{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
IsActive: source.IsActive,
|
||||
IconURL: source.IconURL,
|
||||
ClientSecretConfigured: source.ClientSecretConfigured,
|
||||
})
|
||||
}
|
||||
|
||||
dbSources, err := model.GetActiveAuthSources()
|
||||
if err != nil {
|
||||
return sources
|
||||
return nil
|
||||
}
|
||||
sources := make([]AuthSourceView, 0, len(dbSources))
|
||||
for _, source := range dbSources {
|
||||
sources = append(sources, AuthSourceView{
|
||||
ID: source.ID,
|
||||
@@ -126,9 +108,6 @@ func frontendLoginRedirectURL() string {
|
||||
if config.Config.App.FrontendURL != "" {
|
||||
return strings.TrimRight(config.Config.App.FrontendURL, "/") + "/login"
|
||||
}
|
||||
if config.Config.OAuth2.RedirectURI != "" {
|
||||
return config.Config.OAuth2.RedirectURI
|
||||
}
|
||||
return "/login"
|
||||
}
|
||||
|
||||
@@ -137,24 +116,6 @@ func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL
|
||||
return nil, nil, errors.New("认证源不能为空")
|
||||
}
|
||||
|
||||
if source.Name == "default" {
|
||||
scopes := []string{"profile", "email"}
|
||||
if oidcVerifier != nil {
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
return &oauth2.Config{
|
||||
ClientID: config.Config.OAuth2.ClientID,
|
||||
ClientSecret: config.Config.OAuth2.ClientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: scopes,
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: config.Config.OAuth2.AuthorizationEndpoint,
|
||||
TokenURL: config.Config.OAuth2.TokenEndpoint,
|
||||
AuthStyle: oauth2.AuthStyleAutoDetect,
|
||||
},
|
||||
}, oidcVerifier, nil
|
||||
}
|
||||
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return nil, nil, errors.New("OIDC 认证源必须配置 Discovery URL")
|
||||
}
|
||||
@@ -260,29 +221,6 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
|
||||
userInfo.Name = userInfo.Username
|
||||
}
|
||||
|
||||
if userInfo.Username == "" {
|
||||
client := authConfig.Client(ctx, token)
|
||||
userEndpoint := config.Config.OAuth2.UserEndpoint
|
||||
if source.Name != "default" {
|
||||
userEndpoint = ""
|
||||
}
|
||||
if userEndpoint != "" {
|
||||
resp, httpErr := client.Get(userEndpoint)
|
||||
if httpErr != nil {
|
||||
return nil, httpErr
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
responseData, readErr := io.ReadAll(resp.Body)
|
||||
if readErr != nil {
|
||||
return nil, readErr
|
||||
}
|
||||
if unmarshalErr := json.Unmarshal(responseData, userInfo); unmarshalErr != nil {
|
||||
return nil, unmarshalErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return userInfo, nil
|
||||
}
|
||||
|
||||
@@ -336,10 +274,10 @@ func GetLoginSources(c *gin.Context) {
|
||||
|
||||
// GetLoginURL 获取登录授权地址
|
||||
// @Summary 获取登录授权地址
|
||||
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用默认认证源。
|
||||
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Param source query string false "认证源名称,为空使用默认源"
|
||||
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
||||
// @Failure 400 {object} util.ResponseAny "认证源不存在或未配置"
|
||||
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
||||
@@ -522,16 +460,8 @@ func Callback(c *gin.Context) {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error()))
|
||||
return
|
||||
}
|
||||
user = model.User{
|
||||
Username: username,
|
||||
Nickname: userInfo.Name,
|
||||
AvatarUrl: userInfo.AvatarUrl,
|
||||
TrustLevel: userInfo.TrustLevel,
|
||||
SignKey: util.GenerateUniqueIDSimple(),
|
||||
IsActive: true,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
if err := db.DB(ctx).Create(&user).Error; err != nil {
|
||||
userInfo.Username = username
|
||||
if err := user.CreateUser(db.DB(ctx), userInfo); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -92,6 +92,6 @@ func (h *CleanupUnusedUploadsHandler) Execute(ctx context.Context, payload []byt
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("共处理 %d 个文件,成功删除 %d 个", totalProcessed, totalDeleted)
|
||||
task.AppendLog(ctx, msg)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user