refactor(server): merge authserver implementation

This commit is contained in:
mac
2026-07-15 14:05:22 +08:00
parent 8203beb8cd
commit 6af511ed0b
59 changed files with 176 additions and 497 deletions
-8
View File
@@ -1,8 +0,0 @@
# 业务逻辑层
存放核心业务逻辑,编排 Repository:
- auth.go: 认证服务,LDAP 认证、本地降级认证、JWT 生成
- ldap.go: LDAP 服务,封装 LDAP 连接、查询、认证逻辑
- audit.go: 审计服务,记录登录审计、运维操作审计
- subsystem.go: 子系统服务,子系统查询、OAuth 2.0 SSO URL 生成
- ansible.go: Ansible 调度服务,创建任务、执行 Playbook、推送日志
+60
View File
@@ -0,0 +1,60 @@
package service
import (
"encoding/json"
"github.com/1024XEngineer/xinfra/server/internal/model"
"gorm.io/gorm"
)
type AuditService struct {
db *gorm.DB
}
func NewAuditService(db *gorm.DB) *AuditService {
return &AuditService{db: db}
}
type AuditEntry struct {
RequestID string
ActorUserID uint64
ActorUsername string
ClientIP string
UserAgent string
Action string
ResourceType string
ResourceID string
ScopeType string
ScopeID uint64
BusinessLineID uint64
NamespaceID uint64
ClusterID uint64
EnvironmentID uint64
Decision string
Reason string
Metadata map[string]any
}
func (s *AuditService) Write(entry AuditEntry) {
raw, _ := json.Marshal(entry.Metadata)
_ = s.db.Create(&model.AuditLog{
RequestID: entry.RequestID,
ActorUserID: entry.ActorUserID,
ActorUsername: entry.ActorUsername,
ClientIP: entry.ClientIP,
UserAgent: entry.UserAgent,
Action: entry.Action,
ResourceType: entry.ResourceType,
ResourceID: entry.ResourceID,
ScopeType: entry.ScopeType,
ScopeID: entry.ScopeID,
BusinessLineID: entry.BusinessLineID,
NamespaceID: entry.NamespaceID,
ClusterID: entry.ClusterID,
EnvironmentID: entry.EnvironmentID,
Decision: entry.Decision,
Reason: entry.Reason,
Metadata: string(raw),
}).Error
}
+245
View File
@@ -0,0 +1,245 @@
package service
import (
"errors"
"fmt"
"net/mail"
"regexp"
"strings"
"time"
"github.com/1024XEngineer/xinfra/server/internal/auth"
"github.com/1024XEngineer/xinfra/server/internal/config"
"github.com/1024XEngineer/xinfra/server/internal/model"
"github.com/1024XEngineer/xinfra/server/internal/sso"
"gorm.io/gorm"
)
var (
ErrInvalidCredential = errors.New("invalid username or password")
ErrUserDisabled = errors.New("user disabled")
ErrSAMLSubjectMissing = errors.New("saml subject is missing")
)
type AuthService struct {
cfg config.Config
db *gorm.DB
audit *AuditService
}
type LoginResult struct {
Token string `json:"token"`
TokenID string `json:"token_id"`
ExpiresAt time.Time `json:"expires_at"`
User model.User `json:"user"`
}
func NewAuthService(cfg config.Config, db *gorm.DB, audit *AuditService) *AuthService {
return &AuthService{cfg: cfg, db: db, audit: audit}
}
func (s *AuthService) SAMLLogin(info *sso.SAMLDebugInfo, clientIP, userAgent string) (*LoginResult, error) {
subject := strings.TrimSpace(info.NameID)
email := firstSAMLAttribute(info.Attributes,
"email",
"mail",
"Email",
"EmailAddress",
"emailAddress",
"urn:oid:0.9.2342.19200300.100.1.3",
"urn:oid:1.3.6.1.4.1.5923.1.1.1.6",
)
if email == "" && looksLikeEmail(subject) {
email = subject
}
if subject == "" {
subject = email
}
if subject == "" {
reason := fmt.Sprintf("subject_missing name_id=%q attribute_keys=%v", info.NameID, samlAttributeKeys(info.Attributes))
s.audit.Write(AuditEntry{ClientIP: clientIP, UserAgent: userAgent, Action: "saml.login.failed", Decision: "deny", Reason: reason})
return nil, ErrSAMLSubjectMissing
}
displayName := firstSAMLAttribute(info.Attributes,
"displayName",
"name",
"cn",
"username",
"uid",
"urn:oid:2.5.4.3",
"urn:oid:0.9.2342.19200300.100.1.1",
)
var result *LoginResult
err := s.db.Transaction(func(tx *gorm.DB) error {
user, err := findOrCreateSAMLUser(tx, subject, email, displayName)
if err != nil {
return err
}
if user.Status != "active" {
s.audit.Write(AuditEntry{ActorUserID: user.ID, ActorUsername: user.Username, ClientIP: clientIP, UserAgent: userAgent, Action: "saml.login.failed", Decision: "deny", Reason: "user_disabled"})
return ErrUserDisabled
}
token, tokenID, expiresAt, err := auth.Sign(s.cfg.JWTSecret, s.cfg.JWTIssuer, s.cfg.JWTTTL(), user.ID, user.Username, user.Email, user.IsAdmin)
if err != nil {
return err
}
now := time.Now()
if err := tx.Model(&model.User{}).Where("id = ?", user.ID).Updates(map[string]any{
"last_login_at": now,
"email": user.Email,
"display_name": user.DisplayName,
}).Error; err != nil {
return err
}
if err := tx.Create(&model.AccessToken{
UserID: user.ID,
TokenID: tokenID,
TokenType: "access",
ClientIP: clientIP,
UserAgent: userAgent,
ExpiresAt: expiresAt,
}).Error; err != nil {
return err
}
user.LastLoginAt = &now
result = &LoginResult{Token: token, TokenID: tokenID, ExpiresAt: expiresAt, User: user}
return nil
})
if err != nil {
return nil, err
}
s.audit.Write(AuditEntry{ActorUserID: result.User.ID, ActorUsername: result.User.Username, ClientIP: clientIP, UserAgent: userAgent, Action: "saml.login.success", Decision: "allow"})
return result, nil
}
func findOrCreateSAMLUser(tx *gorm.DB, subject, email, displayName string) (model.User, error) {
username := samlUsername(email, subject)
if displayName == "" {
displayName = username
}
var user model.User
if err := findExistingSAMLUser(tx, &user, username, email, subject); err == nil {
return updateSAMLUser(tx, user, username, email, displayName, subject)
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return model.User{}, err
}
user = model.User{
Username: username,
DisplayName: displayName,
Email: email,
Source: "saml",
ExternalID: subject,
Status: "active",
IsAdmin: false,
}
if err := tx.Create(&user).Error; err != nil {
return user, err
}
return user, nil
}
func findExistingSAMLUser(tx *gorm.DB, user *model.User, username, email, subject string) error {
query := tx.Where("username = ? AND deleted_at IS NULL", username)
if email != "" {
query = query.Or("email = ? AND deleted_at IS NULL", email)
}
if subject != "" {
query = query.Or("source = ? AND external_id = ? AND deleted_at IS NULL", "saml", subject)
}
return query.First(user).Error
}
func updateSAMLUser(tx *gorm.DB, user model.User, username, email, displayName, subject string) (model.User, error) {
updates := map[string]any{
"source": "saml",
"external_id": subject,
}
user.Source = "saml"
user.ExternalID = subject
if email != "" && user.Email != email {
updates["email"] = email
user.Email = email
}
if displayName != "" && user.DisplayName != displayName {
updates["display_name"] = displayName
user.DisplayName = displayName
}
if username != "" && user.Username != username {
var count int64
if err := tx.Model(&model.User{}).Where("username = ? AND id <> ? AND deleted_at IS NULL", username, user.ID).Count(&count).Error; err != nil {
return user, err
}
if count == 0 {
updates["username"] = username
user.Username = username
}
}
if err := tx.Model(&model.User{}).Where("id = ?", user.ID).Updates(updates).Error; err != nil {
return user, err
}
return user, nil
}
func firstSAMLAttribute(attrs map[string][]string, names ...string) string {
for _, name := range names {
for _, value := range attrs[name] {
value = strings.TrimSpace(value)
if value != "" {
return value
}
}
}
for key, values := range attrs {
lower := strings.ToLower(key)
if strings.Contains(lower, "email") || strings.Contains(lower, "mail") {
for _, value := range values {
value = strings.TrimSpace(value)
if value != "" {
return value
}
}
}
}
return ""
}
func looksLikeEmail(value string) bool {
if value == "" {
return false
}
_, err := mail.ParseAddress(value)
return err == nil && strings.Contains(value, "@")
}
var usernameCleaner = regexp.MustCompile(`[^a-zA-Z0-9@._-]+`)
func samlUsername(email, subject string) string {
base := strings.TrimSpace(email)
if base == "" {
base = strings.TrimSpace(subject)
}
base = usernameCleaner.ReplaceAllString(base, "-")
base = strings.Trim(base, ".-_")
if base == "" {
return "saml-user"
}
return base
}
func samlAttributeKeys(attrs map[string][]string) []string {
keys := make([]string, 0, len(attrs))
for key := range attrs {
keys = append(keys, key)
}
return keys
}
+172
View File
@@ -0,0 +1,172 @@
package service
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
"github.com/1024XEngineer/xinfra/server/internal/config"
"github.com/1024XEngineer/xinfra/server/internal/model"
"gorm.io/gorm"
)
var (
ErrWayenNotConfigured = errors.New("wayen login is not configured")
ErrWayenEmailMissing = errors.New("email is missing in token")
ErrWayenCredentialNotFound = errors.New("wayen credential not found")
ErrWayenLoginFailed = errors.New("wayen login failed")
)
type WayenService struct {
cfg config.Config
db *gorm.DB
client *http.Client
}
type WayenLoginResult struct {
TargetURL string
Cookies []*http.Cookie
}
func NewWayenService(cfg config.Config, db *gorm.DB) *WayenService {
return &WayenService{
cfg: cfg,
db: db,
client: &http.Client{
Timeout: 10 * time.Second,
CheckRedirect: func(req *http.Request, via []*http.Request) error {
return http.ErrUseLastResponse
},
},
}
}
func (s *WayenService) Login(email, username string) (*WayenLoginResult, error) {
email = strings.TrimSpace(email)
if email == "" {
return nil, ErrWayenEmailMissing
}
if strings.TrimSpace(s.cfg.OAuthRedirectURI) != "" && strings.TrimSpace(s.cfg.WayenTargetURL) != "" {
target, err := s.oauthLoginURL(s.cfg.OAuthRedirectURI, s.cfg.WayenTargetURL)
if err != nil {
return nil, err
}
return &WayenLoginResult{TargetURL: target}, nil
}
if strings.TrimSpace(s.cfg.WayenLoginURL) == "" || strings.TrimSpace(s.cfg.WayenTargetURL) == "" {
return nil, ErrWayenNotConfigured
}
var credential model.WayenCredential
if err := s.db.Where("email = ? AND enabled = ?", email, true).First(&credential).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrWayenCredentialNotFound
}
return nil, err
}
loginName := email
if strings.EqualFold(s.cfg.WayenLoginValue, "username") && strings.TrimSpace(username) != "" {
loginName = strings.TrimSpace(username)
}
loginURL, body, contentType, err := s.loginRequest(s.cfg.WayenLoginURL, loginName, credential.Password)
if err != nil {
return nil, err
}
req, err := http.NewRequest(http.MethodPost, loginURL, body)
if err != nil {
return nil, err
}
if contentType != "" {
req.Header.Set("Content-Type", contentType)
}
req.Header.Set("Accept", "application/json, text/html;q=0.9, */*;q=0.8")
resp, err := s.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20))
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusBadRequest {
return nil, fmt.Errorf("%w: status %d", ErrWayenLoginFailed, resp.StatusCode)
}
return &WayenLoginResult{
TargetURL: s.cfg.WayenTargetURL,
Cookies: resp.Cookies(),
}, nil
}
func (s *WayenService) oauthLoginURL(redirectURI, targetURL string) (string, error) {
parsed, err := url.Parse(strings.TrimSpace(redirectURI))
if err != nil {
return "", err
}
next, err := url.Parse(strings.TrimSpace(targetURL))
if err != nil {
return "", err
}
if next.Path == "" || next.Path == "/" {
next.Path = "/sign-in"
}
values := next.Query()
values.Set("ref", defaultConfigValue(s.cfg.WayenOAuthRef, "/portal/namespace/1/app"))
next.RawQuery = values.Encode()
query := parsed.Query()
query.Set("next", next.String())
parsed.RawQuery = query.Encode()
return parsed.String(), nil
}
func (s *WayenService) loginRequest(loginURL, email, password string) (string, io.Reader, string, error) {
usernameKey := defaultConfigValue(s.cfg.WayenUsernameKey, "email")
passwordKey := defaultConfigValue(s.cfg.WayenPasswordKey, "password")
if strings.EqualFold(s.cfg.WayenLoginFormat, "query") {
parsed, err := url.Parse(loginURL)
if err != nil {
return "", nil, "", err
}
values := parsed.Query()
values.Set(usernameKey, email)
values.Set(passwordKey, password)
parsed.RawQuery = values.Encode()
return parsed.String(), nil, "", nil
}
if strings.EqualFold(s.cfg.WayenLoginFormat, "json") {
payload := map[string]string{
usernameKey: email,
passwordKey: password,
}
data, err := json.Marshal(payload)
if err != nil {
return "", nil, "", err
}
return loginURL, bytes.NewReader(data), "application/json", nil
}
values := url.Values{}
values.Set(usernameKey, email)
values.Set(passwordKey, password)
return loginURL, strings.NewReader(values.Encode()), "application/x-www-form-urlencoded", nil
}
func defaultConfigValue(value, fallback string) string {
value = strings.TrimSpace(value)
if value == "" {
return fallback
}
return value
}