fix: 完成 broker 复审 B-03 至 B-12
每连接异步下发与背压、写出后断开、校验当前连接与订阅、生命周期串行、登录条件更新、闲置按在线计、认证超时并发与 Shutdown 0x8B。
This commit is contained in:
+118
-10
@@ -32,6 +32,9 @@ type Login struct {
|
||||
// 内存中的 session_used_at(毫秒)与上次落库时间。
|
||||
usedAt map[string]int64
|
||||
lastFlush map[string]int64
|
||||
|
||||
verifyMu sync.Mutex
|
||||
verifySem map[string]chan struct{}
|
||||
}
|
||||
|
||||
// LoginOptions 装配 Login。
|
||||
@@ -67,9 +70,15 @@ func NewLogin(opts LoginOptions) *Login {
|
||||
Now: now,
|
||||
usedAt: make(map[string]int64),
|
||||
lastFlush: make(map[string]int64),
|
||||
verifySem: make(map[string]chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
const (
|
||||
authTimeout = 30 * time.Second
|
||||
verifyPerEndpoint = 2
|
||||
)
|
||||
|
||||
// Authenticate 按 DEVELOPMENT 第 5 节校验;内部故障返回 error。
|
||||
func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) {
|
||||
if l == nil || l.DB == nil {
|
||||
@@ -78,6 +87,11 @@ func (l *Login) Authenticate(ctx context.Context, endpointID string, password []
|
||||
if endpointID == "" {
|
||||
return AuthResult{OK: false}, nil
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(ctx, authTimeout)
|
||||
defer cancel()
|
||||
|
||||
row, err := l.loadEndpoint(ctx, endpointID)
|
||||
if err != nil {
|
||||
@@ -107,6 +121,8 @@ type endpointAuthRow struct {
|
||||
loginHash string
|
||||
sessionHash []byte // 原始 32 字节;无令牌时 nil
|
||||
sessionUsedAt int64 // 毫秒;无则 0
|
||||
onlineSince int64
|
||||
offlineSince int64
|
||||
}
|
||||
|
||||
func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) {
|
||||
@@ -115,10 +131,12 @@ func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, e
|
||||
enabled int
|
||||
sessHex sql.NullString
|
||||
usedAt sql.NullInt64
|
||||
online sql.NullInt64
|
||||
offline sql.NullInt64
|
||||
)
|
||||
err := l.DB.Read.QueryRowContext(ctx, `
|
||||
SELECT login_hash, enabled, session_hash, session_used_at
|
||||
FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt)
|
||||
SELECT login_hash, enabled, session_hash, session_used_at, online_since, offline_since
|
||||
FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt, &online, &offline)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return endpointAuthRow{}, ErrEndpointNotFound
|
||||
@@ -132,6 +150,12 @@ FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt)
|
||||
if usedAt.Valid {
|
||||
row.sessionUsedAt = usedAt.Int64
|
||||
}
|
||||
if online.Valid {
|
||||
row.onlineSince = online.Int64
|
||||
}
|
||||
if offline.Valid {
|
||||
row.offlineSince = offline.Int64
|
||||
}
|
||||
if sessHex.Valid && sessHex.String != "" {
|
||||
raw, decErr := hex.DecodeString(sessHex.String)
|
||||
if decErr != nil || len(raw) != 32 {
|
||||
@@ -160,11 +184,8 @@ func (l *Login) authSession(ctx context.Context, endpointID, token string, row e
|
||||
usedAt = mem
|
||||
}
|
||||
l.usedMu.Unlock()
|
||||
if l.IdleDays > 0 {
|
||||
idle := time.Duration(l.IdleDays) * 24 * time.Hour
|
||||
if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle {
|
||||
return false, nil
|
||||
}
|
||||
if l.IdleDays > 0 && !sessionIdleOK(now, usedAt, row.onlineSince, row.offlineSince, l.IdleDays) {
|
||||
return false, nil
|
||||
}
|
||||
if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil {
|
||||
return false, err
|
||||
@@ -204,7 +225,11 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP
|
||||
if l.Pool == nil {
|
||||
return false, "", errors.New("broker: password pool not configured")
|
||||
}
|
||||
if err := l.acquireVerify(ctx, endpointID); err != nil {
|
||||
return false, "", err
|
||||
}
|
||||
match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash)
|
||||
l.releaseVerify(endpointID)
|
||||
if verErr != nil {
|
||||
return false, "", verErr
|
||||
}
|
||||
@@ -220,12 +245,28 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP
|
||||
}
|
||||
nowMs := l.Now().UnixMilli()
|
||||
hashHex := hex.EncodeToString(hash)
|
||||
var oldHex any
|
||||
if len(row.sessionHash) == 0 {
|
||||
oldHex = ""
|
||||
} else {
|
||||
oldHex = hex.EncodeToString(row.sessionHash)
|
||||
}
|
||||
writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`
|
||||
res, e := tx.Exec(`
|
||||
UPDATE endpoints
|
||||
SET session_hash = ?, session_issued_at = ?, session_used_at = ?
|
||||
WHERE id = ?`, hashHex, nowMs, nowMs, endpointID)
|
||||
return e
|
||||
WHERE id = ? AND COALESCE(session_hash, '') = ?`, hashHex, nowMs, nowMs, endpointID, oldHex)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, nErr := res.RowsAffected()
|
||||
if nErr != nil {
|
||||
return nErr
|
||||
}
|
||||
if n == 0 {
|
||||
return ErrSessionWriteConflict
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if writeErr != nil {
|
||||
return false, "", writeErr
|
||||
@@ -288,6 +329,73 @@ func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, e
|
||||
return hex.DecodeString(sessHex.String)
|
||||
}
|
||||
|
||||
// TokenMatchesDB 握手时重读:明文令牌是否仍对应库中当前哈希。
|
||||
func (l *Login) TokenMatchesDB(ctx context.Context, endpointID, token string) (bool, error) {
|
||||
if l == nil || l.DB == nil || token == "" {
|
||||
return false, nil
|
||||
}
|
||||
got := l.Tokens.HashToken(token)
|
||||
dbHash, err := l.SessionHashOf(ctx, endpointID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if len(dbHash) == 0 {
|
||||
return false, nil
|
||||
}
|
||||
return auth.EqualHash(got, dbHash), nil
|
||||
}
|
||||
|
||||
func sessionIdleOK(now time.Time, usedAt, onlineSince, offlineSince int64, idleDays int) bool {
|
||||
if idleDays <= 0 {
|
||||
return true
|
||||
}
|
||||
online := onlineSince > 0 && onlineSince >= offlineSince
|
||||
if online {
|
||||
return true
|
||||
}
|
||||
activity := usedAt
|
||||
if onlineSince > activity {
|
||||
activity = onlineSince
|
||||
}
|
||||
if offlineSince > activity {
|
||||
activity = offlineSince
|
||||
}
|
||||
if activity <= 0 {
|
||||
return false
|
||||
}
|
||||
idle := time.Duration(idleDays) * 24 * time.Hour
|
||||
return now.Sub(time.UnixMilli(activity)) <= idle
|
||||
}
|
||||
|
||||
func (l *Login) acquireVerify(ctx context.Context, endpointID string) error {
|
||||
l.verifyMu.Lock()
|
||||
sem := l.verifySem[endpointID]
|
||||
if sem == nil {
|
||||
sem = make(chan struct{}, verifyPerEndpoint)
|
||||
l.verifySem[endpointID] = sem
|
||||
}
|
||||
l.verifyMu.Unlock()
|
||||
select {
|
||||
case sem <- struct{}{}:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (l *Login) releaseVerify(endpointID string) {
|
||||
l.verifyMu.Lock()
|
||||
sem := l.verifySem[endpointID]
|
||||
l.verifyMu.Unlock()
|
||||
if sem == nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-sem:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// LooksLikeSessionToken 暴露给测试。
|
||||
func (l *Login) LooksLikeSessionToken(s string) bool {
|
||||
return strings.HasPrefix(s, "nst_")
|
||||
|
||||
Reference in New Issue
Block a user