diff --git a/config/config.example.yaml b/config/config.example.yaml
index fb044f6..4ebc8d1 100644
--- a/config/config.example.yaml
+++ b/config/config.example.yaml
@@ -58,6 +58,22 @@ mcp:
timeoutMs: 15000
headers: {}
+oidc:
+ # 是否启用 OIDC 登录;本系统作为 OIDC Client。
+ enabled: false
+ # OIDC Provider issuer,例如 https://idp.example.com/realms/demo。
+ issuer: ""
+ clientId: ""
+ clientSecret: ""
+ # 必须与 OIDC Provider 后台登记的 redirect_uri 完全一致。
+ redirectUrl: "http://127.0.0.1:8083/api/auth/oidc_callback"
+ # OIDC 登录 state 签名密钥;留空时退回使用 clientSecret。
+ stateSecret: ""
+ scopes:
+ - openid
+ - profile
+ - email
+
wxWork:
# 是否启用企业微信能力。
# false 时不会初始化企业微信 SDK,登录、客服回调、应用通知都不可用。
diff --git a/go.mod b/go.mod
index 0720153..89462af 100644
--- a/go.mod
+++ b/go.mod
@@ -47,6 +47,7 @@ require (
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/cloudwego/base64x v0.1.6 // indirect
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.15 // indirect
+ github.com/coreos/go-oidc/v3 v3.18.0 // indirect
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/evanphx/json-patch v0.5.2 // indirect
@@ -54,6 +55,7 @@ require (
github.com/gabriel-vasile/mimetype v1.4.12 // indirect
github.com/gin-contrib/sse v1.1.0 // indirect
github.com/glebarez/go-sqlite v1.21.2 // indirect
+ github.com/go-jose/go-jose/v4 v4.1.4 // indirect
github.com/go-playground/locales v0.14.1 // indirect
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-redis/redis/v8 v8.11.5 // indirect
diff --git a/go.sum b/go.sum
index d091c88..bcff2a9 100644
--- a/go.sum
+++ b/go.sum
@@ -42,6 +42,8 @@ github.com/cloudwego/eino-ext/components/model/openai v0.1.11 h1:juf9kECfmxJBA0r
github.com/cloudwego/eino-ext/components/model/openai v0.1.11/go.mod h1:DBk44Dq1mhuoAacdUzzhZhSGeeBECDI2rIZnJFeVZoE=
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.15 h1:LbdSG9+qWzzp9RFW6dSFkaUW171JvCoYn/K63zX6dQE=
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.15/go.mod h1:p+l0zBB0GjjX8HTlbTs3g3KfUFwZC11bsCGZOXW/3L0=
+github.com/coreos/go-oidc/v3 v3.18.0 h1:V9orjXynvu5wiC9SemFTWnG4F45v403aIcjWo0d41+A=
+github.com/coreos/go-oidc/v3 v3.18.0/go.mod h1:DYCf24+ncYi+XkIH97GY1+dqoRlbaSI26KVTCI9SrY4=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
@@ -74,6 +76,8 @@ github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GM
github.com/glebarez/sqlite v1.11.0/go.mod h1:h8/o8j5wiAsqSPoWELDUdJXhjAhsVliSn7bWZjOhrgQ=
github.com/go-check/check v0.0.0-20180628173108-788fd7840127 h1:0gkP6mzaMqkmpcJYCFOLkIBwI7xFExG03bbkOkCvUPI=
github.com/go-check/check v0.0.0-20180628173108-788fd7840127/go.mod h1:9ES+weclKsC9YodN5RgxqK/VD9HM9JsCSh7rNhMZE98=
+github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
+github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
diff --git a/internal/bootstrap/init.go b/internal/bootstrap/init.go
index 68f3781..c557032 100644
--- a/internal/bootstrap/init.go
+++ b/internal/bootstrap/init.go
@@ -1,7 +1,9 @@
package bootstrap
import (
+ "context"
"cs-agent/internal/ai/rag/vectordb"
+ "cs-agent/internal/oidcclient"
"cs-agent/internal/pkg/config"
"cs-agent/internal/pkg/logx"
"cs-agent/internal/services/cronx"
@@ -42,5 +44,9 @@ func Init(configPath string) error {
cronx.Init()
wxwork.Init()
+ if err := oidcclient.Init(context.Background()); err != nil {
+ slog.Error("init oidc failed", "error", err)
+ return err
+ }
return nil
}
diff --git a/internal/bootstrap/routes.go b/internal/bootstrap/routes.go
index 4bf9778..683c06c 100644
--- a/internal/bootstrap/routes.go
+++ b/internal/bootstrap/routes.go
@@ -16,6 +16,9 @@ func registerApiAuthRoutes(group *gin.RouterGroup) {
group.POST("/wxwork_exchange", api.WxWorkExchange)
group.GET("/wxwork_login", api.WxWorkLogin)
group.GET("/wxwork_qr_login", api.WxWorkQRLogin)
+ group.GET("/oidc_callback", api.OIDCCallback)
+ group.POST("/oidc_exchange", api.OIDCExchange)
+ group.GET("/oidc_login", api.OIDCLogin)
}
func registerApiChannelRoutes(group *gin.RouterGroup) {
diff --git a/internal/bootstrap/server_route_test.go b/internal/bootstrap/server_route_test.go
index 6e5d6f6..474724b 100644
--- a/internal/bootstrap/server_route_test.go
+++ b/internal/bootstrap/server_route_test.go
@@ -31,6 +31,9 @@ func TestNewServerRegistersGinRoutes(t *testing.T) {
expected := []string{
http.MethodPost + " /api/auth/login",
+ http.MethodGet + " /api/auth/oidc_login",
+ http.MethodGet + " /api/auth/oidc_callback",
+ http.MethodPost + " /api/auth/oidc_exchange",
http.MethodGet + " /api/auth/profile",
http.MethodGet + " /api/dashboard/user/list",
http.MethodGet + " /api/dashboard/user/:id",
diff --git a/internal/handlers/api/auth_handler.go b/internal/handlers/api/auth_handler.go
index c53576d..a552dc8 100644
--- a/internal/handlers/api/auth_handler.go
+++ b/internal/handlers/api/auth_handler.go
@@ -77,6 +77,46 @@ func WxWorkExchange(ctx *gin.Context) {
httpx.WriteJSON(ctx, ret)
}
+func OIDCLogin(ctx *gin.Context) {
+ loginURL, err := services.OIDCLoginService.BuildOIDCLoginURL(ctx.Query("next"))
+ if err != nil {
+ ctx.Redirect(http.StatusFound, "/dashboard/login?oidcError="+url.QueryEscape(loginErrorMessage(err.Error())))
+ return
+ }
+ ctx.Redirect(http.StatusFound, loginURL)
+}
+
+func OIDCCallback(ctx *gin.Context) {
+ cfg := config.Current()
+ ticket, next, err := services.OIDCLoginService.LoginByOIDC(
+ ctx.Request.Context(),
+ ctx.Query("code"),
+ ctx.Query("state"),
+ cfg.Auth,
+ ctx.ClientIP(),
+ ctx.GetHeader("User-Agent"),
+ )
+ if err != nil {
+ ctx.Redirect(http.StatusFound, "/dashboard/login?oidcError="+url.QueryEscape(loginErrorMessage(err.Error())))
+ return
+ }
+ ctx.Redirect(http.StatusFound, "/dashboard/login/oidc/callback?ticket="+url.QueryEscape(ticket)+"&next="+url.QueryEscape(next))
+}
+
+func OIDCExchange(ctx *gin.Context) {
+ req := request.OIDCExchangeRequest{}
+ if err := params.ReadJSON(ctx, &req); err != nil {
+ httpx.WriteJSON(ctx, err)
+ return
+ }
+ ret, err := services.OIDCLoginService.ExchangeOIDCLoginTicket(req.Ticket)
+ if err != nil {
+ httpx.WriteJSON(ctx, err)
+ return
+ }
+ httpx.WriteJSON(ctx, ret)
+}
+
func Logout(ctx *gin.Context) {
if err := services.AuthService.Logout(ctx.GetHeader("Authorization")); err != nil {
httpx.WriteJSON(ctx, err)
@@ -95,6 +135,10 @@ func Profile(ctx *gin.Context) {
}
func wxWorkErrorMessage(message string) string {
+ return loginErrorMessage(message)
+}
+
+func loginErrorMessage(message string) string {
if idx := strings.Index(message, ": "); idx >= 0 {
message = message[idx+2:]
}
diff --git a/internal/oidcclient/oidcclient.go b/internal/oidcclient/oidcclient.go
new file mode 100644
index 0000000..6b12889
--- /dev/null
+++ b/internal/oidcclient/oidcclient.go
@@ -0,0 +1,332 @@
+package oidcclient
+
+import (
+ "context"
+ "crypto/hmac"
+ "crypto/rand"
+ "crypto/sha256"
+ "cs-agent/internal/pkg/config"
+ "cs-agent/internal/pkg/dto/response"
+ "cs-agent/internal/pkg/errorsx"
+ "encoding/base64"
+ "encoding/hex"
+ "encoding/json"
+ "fmt"
+ "strings"
+ "sync"
+ "time"
+
+ gooidc "github.com/coreos/go-oidc/v3/oidc"
+ "golang.org/x/oauth2"
+)
+
+const (
+ StateTTL = 5 * time.Minute
+ LoginTicketTTL = 1 * time.Minute
+ defaultLoginNextPath = "/dashboard"
+)
+
+var (
+ oidcCfg config.OIDCConfig
+ provider *gooidc.Provider
+ oauthConfig *oauth2.Config
+ idTokenVerifier *gooidc.IDTokenVerifier
+ loginTicketStore sync.Map
+)
+
+type Profile struct {
+ Subject string `json:"sub"`
+ Email string `json:"email,omitempty"`
+ PreferredUsername string `json:"preferred_username,omitempty"`
+ Name string `json:"name,omitempty"`
+ Picture string `json:"picture,omitempty"`
+ RawProfile string `json:"-"`
+}
+
+type statePayload struct {
+ Next string `json:"next"`
+ Nonce string `json:"nonce"`
+ ExpiredAt int64 `json:"expiredAt"`
+}
+
+type loginTicket struct {
+ Response *response.LoginResponse
+ ExpiredAt time.Time
+}
+
+func Init(ctx context.Context) error {
+ provider = nil
+ oauthConfig = nil
+ idTokenVerifier = nil
+ oidcCfg = config.OIDCConfig{}
+
+ cfg := config.Current().OIDC
+ if !cfg.Enabled {
+ return nil
+ }
+ oidcCfg = cfg
+ if strings.TrimSpace(cfg.Issuer) == "" {
+ return fmt.Errorf("OIDC issuer 未配置")
+ }
+ if strings.TrimSpace(cfg.ClientID) == "" {
+ return fmt.Errorf("OIDC clientId 未配置")
+ }
+ if strings.TrimSpace(cfg.ClientSecret) == "" {
+ return fmt.Errorf("OIDC clientSecret 未配置")
+ }
+ if strings.TrimSpace(cfg.RedirectURL) == "" {
+ return fmt.Errorf("OIDC redirectUrl 未配置")
+ }
+
+ p, err := gooidc.NewProvider(ctx, strings.TrimSpace(cfg.Issuer))
+ if err != nil {
+ return err
+ }
+ scopes := cfg.Scopes
+ if len(scopes) == 0 {
+ scopes = []string{gooidc.ScopeOpenID, "profile", "email"}
+ }
+ provider = p
+ oauthConfig = &oauth2.Config{
+ ClientID: strings.TrimSpace(cfg.ClientID),
+ ClientSecret: strings.TrimSpace(cfg.ClientSecret),
+ Endpoint: p.Endpoint(),
+ RedirectURL: strings.TrimSpace(cfg.RedirectURL),
+ Scopes: scopes,
+ }
+ idTokenVerifier = p.Verifier(&gooidc.Config{ClientID: strings.TrimSpace(cfg.ClientID)})
+ return nil
+}
+
+func Enabled() bool {
+ return oidcCfg.Enabled && provider != nil && oauthConfig != nil && idTokenVerifier != nil
+}
+
+func BuildAuthCodeURL(next string) (string, error) {
+ if !Enabled() {
+ return "", errorsx.BusinessError(1, "OIDC 登录未启用")
+ }
+ state, err := CreateState(next)
+ if err != nil {
+ return "", err
+ }
+ return oauthConfig.AuthCodeURL(state), nil
+}
+
+func ExchangeCode(ctx context.Context, code string) (*Profile, error) {
+ if !Enabled() {
+ return nil, errorsx.BusinessError(1, "OIDC 登录未启用")
+ }
+ code = strings.TrimSpace(code)
+ if code == "" {
+ return nil, errorsx.InvalidParam("OIDC 授权 code 不能为空")
+ }
+ token, err := oauthConfig.Exchange(ctx, code)
+ if err != nil {
+ return nil, err
+ }
+ rawIDToken, ok := token.Extra("id_token").(string)
+ if !ok || strings.TrimSpace(rawIDToken) == "" {
+ return nil, errorsx.Unauthorized("OIDC id_token 不存在")
+ }
+ idToken, err := idTokenVerifier.Verify(ctx, rawIDToken)
+ if err != nil {
+ return nil, err
+ }
+ profile, err := profileFromIDToken(idToken)
+ if err != nil {
+ return nil, err
+ }
+ userInfo, err := provider.UserInfo(ctx, oauth2.StaticTokenSource(token))
+ if err == nil && userInfo != nil && strings.TrimSpace(userInfo.Subject) == profile.Subject {
+ if mergedProfile, mergeErr := profileFromUserInfo(userInfo, profile); mergeErr == nil {
+ profile = mergedProfile
+ }
+ }
+ return profile, nil
+}
+
+func CreateState(next string) (string, error) {
+ secret := stateSecret()
+ if secret == "" {
+ return "", errorsx.BusinessError(2, "OIDC 登录密钥未配置")
+ }
+ nonce, err := randomToken("os_")
+ if err != nil {
+ return "", err
+ }
+ payload := statePayload{
+ Next: sanitizeNextPath(next),
+ Nonce: nonce,
+ ExpiredAt: time.Now().Add(StateTTL).Unix(),
+ }
+ body, err := json.Marshal(payload)
+ if err != nil {
+ return "", err
+ }
+ encoded := base64.RawURLEncoding.EncodeToString(body)
+ return encoded + "." + signState(encoded, secret), nil
+}
+
+func ParseState(state string) (string, error) {
+ secret := stateSecret()
+ if secret == "" {
+ return "", errorsx.Unauthorized("OIDC 登录状态无效或已过期")
+ }
+ parts := strings.Split(strings.TrimSpace(state), ".")
+ if len(parts) != 2 {
+ return "", errorsx.Unauthorized("OIDC 登录状态无效或已过期")
+ }
+ if !hmac.Equal([]byte(parts[1]), []byte(signState(parts[0], secret))) {
+ return "", errorsx.Unauthorized("OIDC 登录状态无效或已过期")
+ }
+ body, err := base64.RawURLEncoding.DecodeString(parts[0])
+ if err != nil {
+ return "", errorsx.Unauthorized("OIDC 登录状态无效或已过期")
+ }
+ payload := statePayload{}
+ if err = json.Unmarshal(body, &payload); err != nil {
+ return "", errorsx.Unauthorized("OIDC 登录状态无效或已过期")
+ }
+ if payload.ExpiredAt <= time.Now().Unix() {
+ return "", errorsx.Unauthorized("OIDC 登录状态无效或已过期")
+ }
+ return sanitizeNextPath(payload.Next), nil
+}
+
+func IssueLoginTicket(loginResp *response.LoginResponse) (string, error) {
+ if loginResp == nil {
+ return "", fmt.Errorf("登录结果不能为空")
+ }
+ ticket, err := randomToken("olt_")
+ if err != nil {
+ return "", err
+ }
+ cleanupExpiredLoginTickets()
+ loginTicketStore.Store(ticket, loginTicket{
+ Response: loginResp,
+ ExpiredAt: time.Now().Add(LoginTicketTTL),
+ })
+ return ticket, nil
+}
+
+func ConsumeLoginTicket(ticket string) (*response.LoginResponse, error) {
+ ticket = strings.TrimSpace(ticket)
+ if ticket == "" {
+ return nil, errorsx.InvalidParam("ticket 不能为空")
+ }
+ value, ok := loginTicketStore.LoadAndDelete(ticket)
+ if !ok {
+ return nil, errorsx.Unauthorized("登录票据无效或已过期")
+ }
+ record, ok := value.(loginTicket)
+ if !ok || record.Response == nil || time.Now().After(record.ExpiredAt) {
+ return nil, errorsx.Unauthorized("登录票据无效或已过期")
+ }
+ return record.Response, nil
+}
+
+func profileFromIDToken(idToken *gooidc.IDToken) (*Profile, error) {
+ var claims map[string]any
+ if err := idToken.Claims(&claims); err != nil {
+ return nil, err
+ }
+ raw, _ := json.Marshal(claims)
+ profile := &Profile{
+ Subject: claimString(claims, "sub"),
+ Email: claimString(claims, "email"),
+ PreferredUsername: claimString(claims, "preferred_username"),
+ Name: claimString(claims, "name"),
+ Picture: claimString(claims, "picture"),
+ RawProfile: string(raw),
+ }
+ if strings.TrimSpace(profile.Subject) == "" {
+ return nil, errorsx.Unauthorized("OIDC 用户标识不存在")
+ }
+ return profile, nil
+}
+
+func profileFromUserInfo(userInfo *gooidc.UserInfo, fallback *Profile) (*Profile, error) {
+ var claims map[string]any
+ if err := userInfo.Claims(&claims); err != nil {
+ return nil, err
+ }
+ raw, _ := json.Marshal(claims)
+ profile := &Profile{
+ Subject: strings.TrimSpace(userInfo.Subject),
+ Email: claimString(claims, "email"),
+ PreferredUsername: claimString(claims, "preferred_username"),
+ Name: claimString(claims, "name"),
+ Picture: claimString(claims, "picture"),
+ RawProfile: string(raw),
+ }
+ if fallback != nil {
+ profile.Email = firstNonEmpty(profile.Email, fallback.Email)
+ profile.PreferredUsername = firstNonEmpty(profile.PreferredUsername, fallback.PreferredUsername)
+ profile.Name = firstNonEmpty(profile.Name, fallback.Name)
+ profile.Picture = firstNonEmpty(profile.Picture, fallback.Picture)
+ }
+ if profile.RawProfile == "" {
+ if fallback != nil {
+ profile.RawProfile = fallback.RawProfile
+ }
+ }
+ if strings.TrimSpace(profile.Subject) == "" {
+ return nil, errorsx.Unauthorized("OIDC 用户标识不存在")
+ }
+ return profile, nil
+}
+
+func firstNonEmpty(values ...string) string {
+ for _, value := range values {
+ if value = strings.TrimSpace(value); value != "" {
+ return value
+ }
+ }
+ return ""
+}
+
+func claimString(claims map[string]any, key string) string {
+ value, _ := claims[key].(string)
+ return strings.TrimSpace(value)
+}
+
+func stateSecret() string {
+ if strings.TrimSpace(oidcCfg.StateSecret) != "" {
+ return strings.TrimSpace(oidcCfg.StateSecret)
+ }
+ return strings.TrimSpace(oidcCfg.ClientSecret)
+}
+
+func signState(content, secret string) string {
+ mac := hmac.New(sha256.New, []byte(secret))
+ _, _ = mac.Write([]byte(content))
+ return hex.EncodeToString(mac.Sum(nil))
+}
+
+func cleanupExpiredLoginTickets() {
+ now := time.Now()
+ loginTicketStore.Range(func(key, value any) bool {
+ record, ok := value.(loginTicket)
+ if !ok || now.After(record.ExpiredAt) {
+ loginTicketStore.Delete(key)
+ }
+ return true
+ })
+}
+
+func sanitizeNextPath(next string) string {
+ next = strings.TrimSpace(next)
+ if next == "" || !strings.HasPrefix(next, "/") || strings.HasPrefix(next, "//") {
+ return defaultLoginNextPath
+ }
+ return next
+}
+
+func randomToken(prefix string) (string, error) {
+ buf := make([]byte, 24)
+ if _, err := rand.Read(buf); err != nil {
+ return "", err
+ }
+ return prefix + hex.EncodeToString(buf), nil
+}
diff --git a/internal/pkg/config/config.go b/internal/pkg/config/config.go
index 1558f34..2755c2f 100644
--- a/internal/pkg/config/config.go
+++ b/internal/pkg/config/config.go
@@ -17,6 +17,7 @@ type Config struct {
VectorDB VectorDBConfig `yaml:"vectorDB"`
MCP MCPConfig `yaml:"mcp"`
WxWork WxWorkConfig `yaml:"wxWork"`
+ OIDC OIDCConfig `yaml:"oidc"`
CustomerSession CustomerSessionConfig `yaml:"customerSession"`
}
@@ -135,6 +136,16 @@ type MCPServerConfig struct {
Headers map[string]string `yaml:"headers"`
}
+type OIDCConfig struct {
+ Enabled bool `yaml:"enabled"`
+ Issuer string `yaml:"issuer"`
+ ClientID string `yaml:"clientId"`
+ ClientSecret string `yaml:"clientSecret"`
+ RedirectURL string `yaml:"redirectUrl"`
+ StateSecret string `yaml:"stateSecret"`
+ Scopes []string `yaml:"scopes"`
+}
+
// WxWorkConfig 定义企业微信接入配置。
//
// 当前主要用于后台管理台的企业微信登录流程:
diff --git a/internal/pkg/dto/request/auth_request.go b/internal/pkg/dto/request/auth_request.go
index 112bd87..51f1ae5 100644
--- a/internal/pkg/dto/request/auth_request.go
+++ b/internal/pkg/dto/request/auth_request.go
@@ -8,3 +8,7 @@ type LoginRequest struct {
type WxWorkExchangeRequest struct {
Ticket string `json:"ticket"`
}
+
+type OIDCExchangeRequest struct {
+ Ticket string `json:"ticket"`
+}
diff --git a/internal/pkg/enums/enums.go b/internal/pkg/enums/enums.go
index 249be9b..8a6e7f0 100644
--- a/internal/pkg/enums/enums.go
+++ b/internal/pkg/enums/enums.go
@@ -63,13 +63,15 @@ type ThirdProvider string
const (
ThirdProviderWxWork ThirdProvider = "wxwork"
ThirdProviderDingtalk ThirdProvider = "dingtalk"
+ ThirdProviderOIDC ThirdProvider = "oidc"
)
-var ThirdProviderValues = []ThirdProvider{ThirdProviderWxWork, ThirdProviderDingtalk}
+var ThirdProviderValues = []ThirdProvider{ThirdProviderWxWork, ThirdProviderDingtalk, ThirdProviderOIDC}
var thirdProviderLabelMap = map[ThirdProvider]string{
ThirdProviderWxWork: "企业微信",
ThirdProviderDingtalk: "钉钉",
+ ThirdProviderOIDC: "OIDC",
}
func GetThirdProviderLabel(provider ThirdProvider) string {
diff --git a/internal/services/auth_service_test.go b/internal/services/auth_service_test.go
index f30a9b0..da5585e 100644
--- a/internal/services/auth_service_test.go
+++ b/internal/services/auth_service_test.go
@@ -381,6 +381,7 @@ func setupAuthServiceTestDB(t *testing.T) *gorm.DB {
}
if err := db.AutoMigrate(
&models.User{},
+ &models.UserIdentity{},
&models.Role{},
&models.Permission{},
&models.UserRole{},
diff --git a/internal/services/oidc_login_service.go b/internal/services/oidc_login_service.go
new file mode 100644
index 0000000..27fd939
--- /dev/null
+++ b/internal/services/oidc_login_service.go
@@ -0,0 +1,245 @@
+package services
+
+import (
+ "context"
+ "crypto/sha256"
+ "cs-agent/internal/models"
+ "cs-agent/internal/oidcclient"
+ "cs-agent/internal/pkg/config"
+ "cs-agent/internal/pkg/dto/response"
+ "cs-agent/internal/pkg/enums"
+ "cs-agent/internal/pkg/errorsx"
+ "cs-agent/internal/repositories"
+ "encoding/hex"
+ "strconv"
+ "strings"
+ "time"
+ "unicode"
+
+ "github.com/mlogclub/simple/sqls"
+ "gorm.io/gorm"
+)
+
+var OIDCLoginService = newOIDCLoginService()
+
+type oidcLoginService struct {
+}
+
+type oidcLoginProfile = oidcclient.Profile
+
+func newOIDCLoginService() *oidcLoginService {
+ return &oidcLoginService{}
+}
+
+func (s *oidcLoginService) BuildOIDCLoginURL(next string) (string, error) {
+ return oidcclient.BuildAuthCodeURL(next)
+}
+
+func (s *oidcLoginService) LoginByOIDC(ctx context.Context, code, state string, authCfg config.AuthConfig, clientIP, userAgent string) (string, string, error) {
+ next, err := oidcclient.ParseState(state)
+ if err != nil {
+ return "", "", err
+ }
+ profile, err := oidcclient.ExchangeCode(ctx, code)
+ if err != nil {
+ return "", "", err
+ }
+ loginResp, err := s.loginWithOIDCProfile(profile, authCfg, clientIP, userAgent)
+ if err != nil {
+ return "", "", err
+ }
+ ticket, err := oidcclient.IssueLoginTicket(loginResp)
+ if err != nil {
+ return "", "", err
+ }
+ return ticket, next, nil
+}
+
+func (s *oidcLoginService) ExchangeOIDCLoginTicket(ticket string) (*response.LoginResponse, error) {
+ return oidcclient.ConsumeLoginTicket(ticket)
+}
+
+func (s *oidcLoginService) loginWithOIDCProfile(profile *oidcLoginProfile, authCfg config.AuthConfig, clientIP, userAgent string) (*response.LoginResponse, error) {
+ if profile == nil || strings.TrimSpace(profile.Subject) == "" {
+ return nil, errorsx.BusinessError(2, "OIDC 用户信息不存在")
+ }
+
+ var ret *response.LoginResponse
+ err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
+ var (
+ identity = repositories.UserIdentityRepository.GetBy(ctx.Tx, enums.ThirdProviderOIDC, "", profile.Subject)
+ user *models.User
+ err error
+ )
+ if identity == nil {
+ user, identity, err = s.createOIDCUser(ctx, profile)
+ if err != nil {
+ return err
+ }
+ } else {
+ if identity.Status != enums.StatusOk {
+ return errorsx.BusinessError(3, "当前 OIDC 绑定已停用")
+ }
+ user = repositories.UserRepository.Get(ctx.Tx, identity.UserID)
+ if user == nil {
+ return errorsx.BusinessError(4, "OIDC 账号绑定的系统用户不存在")
+ }
+ }
+
+ if user.Status != enums.StatusOk {
+ return errorsx.Unauthorized("当前系统账号已被禁用")
+ }
+
+ if err = repositories.UserRepository.Updates(ctx.Tx, user.ID, map[string]any{
+ "nickname": s.resolveOIDCNickname(user.Nickname, profile),
+ "avatar": s.resolveOIDCAvatar(user.Avatar, profile),
+ "last_login_at": time.Now(),
+ "last_login_ip": clientIP,
+ "update_user_id": user.ID,
+ "update_user_name": user.Username,
+ "updated_at": time.Now(),
+ }); err != nil {
+ return err
+ }
+
+ if err = repositories.UserIdentityRepository.Updates(ctx.Tx, identity.ID, map[string]any{
+ "provider_name": enums.GetThirdProviderLabel(enums.ThirdProviderOIDC),
+ "raw_profile": profile.RawProfile,
+ "last_auth_at": time.Now(),
+ "status": enums.StatusOk,
+ "update_user_id": user.ID,
+ "update_user_name": user.Username,
+ "updated_at": time.Now(),
+ }); err != nil {
+ return err
+ }
+
+ ret, err = AuthService.issueTokens(ctx, user, clientIP, userAgent, authCfg)
+ return err
+ })
+ if err != nil {
+ return nil, err
+ }
+ return ret, nil
+}
+
+func (s *oidcLoginService) createOIDCUser(ctx *sqls.TxContext, profile *oidcLoginProfile) (*models.User, *models.UserIdentity, error) {
+ now := time.Now()
+ email := s.availableEmail(ctx.Tx, profile.Email)
+ username := s.availableUsername(ctx.Tx, profile)
+
+ user := &models.User{
+ Username: username,
+ Nickname: s.resolveOIDCNickname("", profile),
+ Avatar: s.resolveOIDCAvatar("", profile),
+ Email: email,
+ Password: "",
+ PasswordSalt: "",
+ Status: enums.StatusOk,
+ AuditFields: models.AuditFields{
+ CreatedAt: now,
+ CreateUserID: 0,
+ CreateUserName: enums.GetThirdProviderLabel(enums.ThirdProviderOIDC),
+ UpdatedAt: now,
+ UpdateUserID: 0,
+ UpdateUserName: enums.GetThirdProviderLabel(enums.ThirdProviderOIDC),
+ },
+ }
+ if err := repositories.UserRepository.Create(ctx.Tx, user); err != nil {
+ return nil, nil, err
+ }
+
+ identity := &models.UserIdentity{
+ UserID: user.ID,
+ Provider: enums.ThirdProviderOIDC,
+ ProviderUserID: strings.TrimSpace(profile.Subject),
+ ProviderCorpID: "",
+ ProviderName: enums.GetThirdProviderLabel(enums.ThirdProviderOIDC),
+ RawProfile: profile.RawProfile,
+ Status: enums.StatusOk,
+ LastAuthAt: &now,
+ AuditFields: models.AuditFields{
+ CreatedAt: now,
+ CreateUserID: user.ID,
+ CreateUserName: user.Username,
+ UpdatedAt: now,
+ UpdateUserID: user.ID,
+ UpdateUserName: user.Username,
+ },
+ }
+ if err := repositories.UserIdentityRepository.Create(ctx.Tx, identity); err != nil {
+ return nil, nil, err
+ }
+ return user, identity, nil
+}
+
+func (s *oidcLoginService) availableEmail(tx *gorm.DB, email string) *string {
+ email = strings.TrimSpace(strings.ToLower(email))
+ if email == "" || repositories.UserRepository.GetByEmail(tx, email) != nil {
+ return nil
+ }
+ return &email
+}
+
+func (s *oidcLoginService) availableUsername(tx *gorm.DB, profile *oidcLoginProfile) string {
+ for _, candidate := range []string{
+ profile.PreferredUsername,
+ strings.Split(strings.TrimSpace(profile.Email), "@")[0],
+ } {
+ username := normalizeOIDCUsername(candidate)
+ if username != "" && repositories.UserRepository.GetByUsername(tx, username) == nil {
+ return username
+ }
+ }
+ base := "oidc_" + shortSubjectHash(profile.Subject)
+ if repositories.UserRepository.GetByUsername(tx, base) == nil {
+ return base
+ }
+ for i := 1; i < 100; i++ {
+ username := base + "_" + strconv.Itoa(i)
+ if repositories.UserRepository.GetByUsername(tx, username) == nil {
+ return username
+ }
+ }
+ return base + "_" + shortSubjectHash(time.Now().String())
+}
+
+func (s *oidcLoginService) resolveOIDCNickname(current string, profile *oidcLoginProfile) string {
+ if profile != nil {
+ for _, candidate := range []string{profile.Name, profile.PreferredUsername, profile.Email, profile.Subject} {
+ if candidate = strings.TrimSpace(candidate); candidate != "" {
+ return candidate
+ }
+ }
+ }
+ return strings.TrimSpace(current)
+}
+
+func (s *oidcLoginService) resolveOIDCAvatar(current string, profile *oidcLoginProfile) string {
+ if profile != nil {
+ if picture := strings.TrimSpace(profile.Picture); picture != "" {
+ return picture
+ }
+ }
+ return strings.TrimSpace(current)
+}
+
+func normalizeOIDCUsername(value string) string {
+ value = strings.TrimSpace(strings.ToLower(value))
+ var b strings.Builder
+ for _, r := range value {
+ if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '_' || r == '-' || r == '.' {
+ b.WriteRune(r)
+ }
+ }
+ ret := strings.Trim(b.String(), "._-")
+ if len(ret) > 100 {
+ ret = ret[:100]
+ }
+ return ret
+}
+
+func shortSubjectHash(subject string) string {
+ sum := sha256.Sum256([]byte(strings.TrimSpace(subject)))
+ return hex.EncodeToString(sum[:])[:16]
+}
diff --git a/internal/services/oidc_login_service_test.go b/internal/services/oidc_login_service_test.go
new file mode 100644
index 0000000..cce2d64
--- /dev/null
+++ b/internal/services/oidc_login_service_test.go
@@ -0,0 +1,96 @@
+package services
+
+import (
+ "strings"
+ "testing"
+
+ "cs-agent/internal/models"
+ "cs-agent/internal/pkg/config"
+ "cs-agent/internal/pkg/enums"
+)
+
+func TestOIDCLoginAutoCreatesSystemUser(t *testing.T) {
+ db := setupAuthServiceTestDB(t)
+ svc := newOIDCLoginService()
+
+ ret, err := svc.loginWithOIDCProfile(&oidcLoginProfile{
+ Subject: "sub-123",
+ Email: "ada@example.com",
+ PreferredUsername: "ada",
+ Name: "Ada Lovelace",
+ Picture: "https://example.com/ada.png",
+ RawProfile: `{"sub":"sub-123"}`,
+ }, config.AuthConfig{TokenTTLHours: 2}, "127.0.0.1", "go-test")
+ if err != nil {
+ t.Fatalf("loginWithOIDCProfile() error = %v", err)
+ }
+ if ret == nil || !strings.HasPrefix(ret.AccessToken, "ak_") {
+ t.Fatalf("expected ak_ access token, got %+v", ret)
+ }
+
+ var user models.User
+ if err := db.Take(&user, "username = ?", "ada").Error; err != nil {
+ t.Fatalf("expected OIDC user to be created: %v", err)
+ }
+ if user.Nickname != "Ada Lovelace" || user.Avatar != "https://example.com/ada.png" {
+ t.Fatalf("unexpected created user profile: %+v", user)
+ }
+ if user.Email == nil || *user.Email != "ada@example.com" {
+ t.Fatalf("expected email to be stored, got %+v", user.Email)
+ }
+ if user.Password != "" {
+ t.Fatalf("expected OIDC-created user password to be empty, got %q", user.Password)
+ }
+
+ var identity models.UserIdentity
+ if err := db.Take(&identity, "provider = ? AND provider_user_id = ?", enums.ThirdProviderOIDC, "sub-123").Error; err != nil {
+ t.Fatalf("expected OIDC identity to be created: %v", err)
+ }
+ if identity.UserID != user.ID || identity.ProviderName != "OIDC" || identity.Status != enums.StatusOk {
+ t.Fatalf("unexpected OIDC identity: %+v", identity)
+ }
+
+ var sessions []models.LoginSession
+ if err := db.Find(&sessions).Error; err != nil {
+ t.Fatalf("query login sessions: %v", err)
+ }
+ if len(sessions) != 1 || sessions[0].UserID != user.ID || sessions[0].Token != ret.AccessToken {
+ t.Fatalf("unexpected login sessions: %+v", sessions)
+ }
+}
+
+func TestOIDCLoginReusesExistingIdentity(t *testing.T) {
+ db := setupAuthServiceTestDB(t)
+ user := createAuthTestUser(t, db, "existing", "secret")
+ if err := db.Create(&models.UserIdentity{
+ UserID: user.ID,
+ Provider: enums.ThirdProviderOIDC,
+ ProviderUserID: "sub-123",
+ ProviderName: "OIDC",
+ Status: enums.StatusOk,
+ }).Error; err != nil {
+ t.Fatalf("seed OIDC identity: %v", err)
+ }
+
+ ret, err := newOIDCLoginService().loginWithOIDCProfile(&oidcLoginProfile{
+ Subject: "sub-123",
+ PreferredUsername: "ignored",
+ Name: "Updated Name",
+ Picture: "https://example.com/updated.png",
+ RawProfile: `{"sub":"sub-123"}`,
+ }, config.AuthConfig{TokenTTLHours: 2}, "127.0.0.1", "go-test")
+ if err != nil {
+ t.Fatalf("loginWithOIDCProfile() error = %v", err)
+ }
+ if ret == nil || ret.User == nil || ret.User.ID != user.ID {
+ t.Fatalf("expected existing user login response, got %+v", ret)
+ }
+
+ var count int64
+ if err := db.Model(&models.User{}).Count(&count).Error; err != nil {
+ t.Fatalf("count users: %v", err)
+ }
+ if count != 1 {
+ t.Fatalf("expected existing identity to reuse user, got %d users", count)
+ }
+}
diff --git a/internal/services/wxwork_login_service.go b/internal/services/wxwork_login_service.go
index cb47a95..017bb33 100644
--- a/internal/services/wxwork_login_service.go
+++ b/internal/services/wxwork_login_service.go
@@ -139,6 +139,7 @@ func (s *wxWorkLoginService) createWxWorkUser(ctx *sqls.TxContext, profile *wxwo
username := strings.TrimSpace(profile.UserID)
mobile := strings.TrimSpace(profile.Mobile)
email := strings.TrimSpace(s.firstNonEmpty(profile.Email, profile.BizMail))
+ now := time.Now()
if err := s.checkWxWorkProfile(ctx.Tx, username, mobile, email); err != nil {
return nil, nil, err
@@ -152,10 +153,10 @@ func (s *wxWorkLoginService) createWxWorkUser(ctx *sqls.TxContext, profile *wxwo
PasswordSalt: "",
Status: enums.StatusOk,
AuditFields: models.AuditFields{
- CreatedAt: time.Now(),
+ CreatedAt: now,
CreateUserID: 0,
CreateUserName: enums.GetThirdProviderLabel(enums.ThirdProviderWxWork),
- UpdatedAt: time.Now(),
+ UpdatedAt: now,
UpdateUserID: 0,
UpdateUserName: enums.GetThirdProviderLabel(enums.ThirdProviderWxWork),
},
@@ -172,12 +173,12 @@ func (s *wxWorkLoginService) createWxWorkUser(ctx *sqls.TxContext, profile *wxwo
ProviderName: enums.GetThirdProviderLabel(enums.ThirdProviderWxWork),
RawProfile: jsons.ToJsonStr(profile),
Status: enums.StatusOk,
- LastAuthAt: new(time.Now()),
+ LastAuthAt: &now,
AuditFields: models.AuditFields{
- CreatedAt: time.Now(),
+ CreatedAt: now,
CreateUserID: user.ID,
CreateUserName: user.Username,
- UpdatedAt: time.Now(),
+ UpdatedAt: now,
UpdateUserID: user.ID,
UpdateUserName: user.Username,
},
diff --git a/web/app/dashboard/login/oidc/callback/page.tsx b/web/app/dashboard/login/oidc/callback/page.tsx
new file mode 100644
index 0000000..fa6894d
--- /dev/null
+++ b/web/app/dashboard/login/oidc/callback/page.tsx
@@ -0,0 +1,63 @@
+"use client"
+
+import { useRouter, useSearchParams } from "next/navigation"
+import { Suspense, useEffect, useRef } from "react"
+import { toast } from "sonner"
+
+import { exchangeOIDCTicket } from "@/lib/api/auth"
+
+export default function OIDCLoginCallbackPage() {
+ return (
+
+ 正在校验登录票据并进入系统,请稍候。 +
+