Files
2026-09-20 00:19:08 +08:00

115 lines
3.1 KiB
Go

package auth
import (
"fmt"
"log"
"net/http"
"net/url"
"strings"
"time"
"github.com/google/uuid"
"server/models"
"server/pkg/jwtutil"
)
// BackchannelLogoutEvent OIDC 标准登出事件声明
const BackchannelLogoutEvent = "http://schemas.openid.net/event/backchannel-logout"
// logoutTarget 需要收到登出通知的应用
type logoutTarget struct {
ClientID string
Endpoint string
Sid string
Identity uint64
}
// NotifyBackchannelLogout 单点登出:通知该用户已登录的所有应用。
//
// 流程:查库收集目标应用(同步)→ 异步逐个 POST logout_token。
// 各应用收到后应清除本地会话,否则用户在这边登出了,其他应用仍显示已登录。
//
// 注意:查库必须在请求上下文内同步完成(beego 全局 Ormer 不适合跨 goroutine 使用),
// 异步部分只做 HTTP 通知,不再触碰数据库。
func NotifyBackchannelLogout(identityID uint64) {
targets, err := collectLogoutTargets(identityID)
if err != nil || len(targets) == 0 {
return
}
go func() {
for _, t := range targets {
if err := sendLogoutToken(t); err != nil {
log.Printf("[auth] 单点登出通知失败 client=%s: %v", t.ClientID, err)
}
}
}()
}
// collectLogoutTargets 收集该用户当前活跃会话涉及的应用(按 client_id 去重)
func collectLogoutTargets(identityID uint64) ([]logoutTarget, error) {
var sessions []models.AuthSession
if _, err := models.Orm.QueryTable(new(models.AuthSession)).
Filter("identity_id", identityID).
Filter("revoked", 0).
All(&sessions); err != nil {
return nil, err
}
seen := map[string]bool{}
targets := make([]logoutTarget, 0)
for _, s := range sessions {
if s.ClientID == "" || seen[s.ClientID] {
continue
}
var client models.AuthClient
if err := models.Orm.QueryTable(new(models.AuthClient)).
Filter("client_id", s.ClientID).One(&client); err != nil {
continue
}
if client.BackchannelLogoutURI == nil || strings.TrimSpace(*client.BackchannelLogoutURI) == "" {
continue
}
seen[s.ClientID] = true
targets = append(targets, logoutTarget{
ClientID: s.ClientID,
Endpoint: strings.TrimSpace(*client.BackchannelLogoutURI),
Sid: s.Sid,
Identity: identityID,
})
}
return targets, nil
}
// sendLogoutToken 按 OIDC Back-Channel Logout 规范发送 logout_token
func sendLogoutToken(t logoutTarget) error {
token, err := jwtutil.SignToken(jwtutil.TokenOptions{
Alg: jwtutil.AlgRS256,
UserID: int(t.Identity),
Subject: fmt.Sprintf("%d", t.Identity),
Audience: []string{t.ClientID},
ClientID: t.ClientID,
Sid: t.Sid,
Jti: uuid.NewString(),
Events: map[string]interface{}{BackchannelLogoutEvent: map[string]interface{}{}},
TTL: 5 * time.Minute,
})
if err != nil {
return err
}
form := url.Values{}
form.Set("logout_token", token)
client := &http.Client{Timeout: 5 * time.Second}
resp, err := client.PostForm(t.Endpoint, form)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 300 {
return fmt.Errorf("应用返回状态码 %d", resp.StatusCode)
}
return nil
}