This commit is contained in:
ryan
2026-06-08 20:08:14 +08:00
parent 02f458856d
commit cd3d0c9f82
61 changed files with 2180 additions and 1891 deletions
+7 -3
View File
@@ -23,9 +23,13 @@ import (
)
const (
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
PendingOAuthSourceIDKey = "pending_oauth_source_id"
PendingOAuthExternalIDKey = "pending_oauth_external_id"
PendingOAuthExternalUsernameKey = "pending_oauth_external_username"
PendingOAuthEmailKey = "pending_oauth_email"
)
const (
+12 -3
View File
@@ -18,7 +18,16 @@ limitations under the License.
package oauth
const (
InvalidState = "非法登录请求"
IDTokenVerifyFailed = "ID Token 验证失败"
NonceMismatch = "nonce 不匹配,可能存在重放攻击"
InvalidState = "非法登录请求"
IDTokenVerifyFailed = "ID Token 验证失败"
IDTokenVerifyFailedFormat = "%s: %w"
NonceMismatch = "nonce 不匹配,可能存在重放攻击"
NoActiveAuthSource = "未配置可用认证源"
ServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
AuthSourceRequired = "认证源不能为空"
DiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
UsernameGenerateFailed = "无法生成可用用户名"
UsernameFromSourceFailed = "无法从认证源获取用户名"
AuthSourceDisabled = "认证源未启用"
InvalidExternalAccountBindingID = "绑定记录 ID 无效"
)
+49 -145
View File
@@ -267,7 +267,50 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string)
// -----------------------------------------------------------------------------
// Test Helpers
// -----------------------------------------------------------------------------
func newMockOIDCClient(issuer, clientID, expectedState, sub, username, email, name string) *http.Client {
cleanIssuer := strings.TrimRight(issuer, "/")
return &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
urlStr := req.URL.String()
if req.Method == http.MethodGet && strings.Contains(urlStr, "/.well-known/openid-configuration") {
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"]
}`, cleanIssuer, cleanIssuer+"/oauth2/authorize", cleanIssuer+"/oauth2/token", cleanIssuer+"/oauth2/keys")
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
}
if req.Method == http.MethodGet && (strings.Contains(urlStr, "/keys") || strings.Contains(urlStr, "/jwks")) {
jwksJSON, _ := json.Marshal(testJWKS)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
Header: make(http.Header),
}, nil
}
if req.Method == http.MethodPost && (strings.Contains(urlStr, "/token") || strings.Contains(urlStr, "/access_token")) {
idToken := generateMockIDToken(issuer, sub, clientID, expectedState, username, email, name)
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)),
Header: make(http.Header),
}, nil
}
return nil, fmt.Errorf("unexpected mock request: %s %s", req.Method, req.URL)
},
},
}
}
func setupTestDB(t *testing.T) *gorm.DB {
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
@@ -593,29 +636,7 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
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) {
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() == 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)),
Header: make(http.Header),
}, nil
}
return nil, fmt.Errorf("unexpected outgoing request: %s %s", req.Method, req.URL)
},
},
}
httpMock := newMockOIDCClient(testIssuerURL, testClientID, state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User")
util.SetHTTPClient(httpMock)
router := setupTestRouter(dbConn, mockRedis, httpMock)
@@ -679,27 +700,7 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state2)), payloadValue, OAuthStateCacheKeyExpiration)
// Callback with same username but different external ID (99999)
httpMock2 := &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
}
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)),
}, nil
}
return nil, fmt.Errorf("unexpected outgoing request")
},
},
}
httpMock2 := newMockOIDCClient(testIssuerURL, testClientID, state2, "99999", "test_oauth_user", "another@linux.do", "Another User")
util.SetHTTPClient(httpMock2)
// Create another router for this mock client
@@ -740,28 +741,7 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
})
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state4)), payloadValue4, OAuthStateCacheKeyExpiration)
httpMock4 := &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
}
if req.Method == http.MethodGet && req.URL.String() == testJWKSURL {
return jwksResponse(), nil
}
if req.Method == http.MethodPost && req.URL.String() == testTokenURL {
idToken := generateMockIDToken(testIssuerURL, "77777", testClientID, state4, "need_bind_user", "needbind@linux.do", "Need Bind 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)),
Header: make(http.Header),
}, nil
}
return nil, fmt.Errorf("unexpected request")
},
},
}
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
util.SetHTTPClient(httpMock4)
router4 := setupTestRouter(dbConn, mockRedis, httpMock4)
@@ -824,47 +804,7 @@ func TestCallbackBind(t *testing.T) {
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration)
// Mock OIDC discovery, JWKS, and Token exchange for custom source (GitHub)
httpMock := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
// Handle OIDC Discovery Document
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") {
body := `{
"issuer": "https://github.com",
"authorization_endpoint": "https://github.com/login/oauth/authorize",
"token_endpoint": "https://github.com/login/oauth/access_token",
"jwks_uri": "https://github.com/oauth/keys",
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}`
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
// Handle JWKS Key Set Fetch
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/oauth/keys") {
jwksJSON, _ := json.Marshal(testJWKS)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
}, nil
}
// Handle Token Exchange
if req.Method == http.MethodPost && strings.Contains(req.URL.String(), "/login/oauth/access_token") {
// Generate signed RS256 token matching issuer and aud
idToken := generateMockIDToken("https://github.com", "github_user_123", "gh_client", state, "github_tester", "tester@github.com", "GitHub Tester")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected request: %s", req.URL)
},
},
}
httpMock := newMockOIDCClient("https://github.com", "gh_client", state, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
util.SetHTTPClient(httpMock)
router := setupTestRouter(dbConn, mockRedis, httpMock)
@@ -951,43 +891,7 @@ func TestCallbackBind(t *testing.T) {
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state3)), payloadValue, OAuthStateCacheKeyExpiration)
// Re-sign token for new state (since state serves as OIDC Nonce)
httpMock3 := &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") {
body := `{
"issuer": "https://github.com",
"authorization_endpoint": "https://github.com/login/oauth/authorize",
"token_endpoint": "https://github.com/login/oauth/access_token",
"jwks_uri": "https://github.com/oauth/keys",
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}`
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/oauth/keys") {
jwksJSON, _ := json.Marshal(testJWKS)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(jwksJSON)),
}, nil
}
if req.Method == http.MethodPost && strings.Contains(req.URL.String(), "/login/oauth/access_token") {
idToken := generateMockIDToken("https://github.com", "github_user_123", "gh_client", state3, "github_tester", "tester@github.com", "GitHub Tester")
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","id_token":"%s"}`, idToken)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected request")
},
},
}
httpMock3 := newMockOIDCClient("https://github.com", "gh_client", state3, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
util.SetHTTPClient(httpMock3)
router3 := setupTestRouter(dbConn, mockRedis, httpMock3)
+13 -13
View File
@@ -87,7 +87,7 @@ func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New("未配置可用认证源")
return nil, errors.New(NoActiveAuthSource)
}
return &sources[0], nil
}
@@ -122,18 +122,18 @@ func activeLoginSources() []AuthSourceView {
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" {
return "", errors.New("服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试")
return "", errors.New(ServerAddressMissing)
}
return strings.TrimRight(sc.Value, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New("认证源不能为空")
return nil, nil, errors.New(AuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New("OIDC 认证源必须配置 Discovery URL")
return nil, nil, errors.New(DiscoveryURLRequired)
}
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
@@ -194,7 +194,7 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
}
candidate = fmt.Sprintf("%s-%d", base, i+1)
}
return "", errors.New("无法生成可用用户名")
return "", errors.New(UsernameGenerateFailed)
}
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
@@ -213,7 +213,7 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
if rawIDToken, ok := token.Extra("id_token").(string); ok {
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return nil, fmt.Errorf("%s: %w", IDTokenVerifyFailed, verifyErr)
return nil, fmt.Errorf(IDTokenVerifyFailedFormat, IDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return nil, errors.New(NonceMismatch)
@@ -257,7 +257,7 @@ func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New("无法从认证源获取用户名")
return errors.New(UsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
@@ -359,7 +359,7 @@ func Authorize(c *gin.Context) {
return
}
if !source.IsActive {
c.JSON(http.StatusBadRequest, util.Err("认证源未启用"))
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
@@ -490,10 +490,10 @@ func Callback(c *gin.Context) {
if !registrationEnabled {
// 如果不允许注册,临时记录到 session 并向前端返回 "need_bind" 状态
session := sessions.Default(c)
session.Set("pending_oauth_source_id", source.ID)
session.Set("pending_oauth_external_id", userInfo.Sub)
session.Set("pending_oauth_external_username", userInfo.Username)
session.Set("pending_oauth_email", userInfo.Email)
session.Set(PendingOAuthSourceIDKey, source.ID)
session.Set(PendingOAuthExternalIDKey, userInfo.Sub)
session.Set(PendingOAuthExternalUsernameKey, userInfo.Username)
session.Set(PendingOAuthEmailKey, userInfo.Email)
if err := session.Save(); err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
return
@@ -577,7 +577,7 @@ func DeleteExternalAccount(c *gin.Context) {
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
c.JSON(http.StatusBadRequest, util.Err("绑定记录 ID 无效"))
c.JSON(http.StatusBadRequest, util.Err(InvalidExternalAccountBindingID))
return
}
if err := model.DeleteExternalAccountForUser(id, userID); err != nil {