Files
NixMsg/internal/broker/authn.go
T
Nixevol 8b845843d6 fix: 完成 broker 复审 B-03 至 B-12
每连接异步下发与背压、写出后断开、校验当前连接与订阅、生命周期串行、登录条件更新、闲置按在线计、认证超时并发与 Shutdown 0x8B。
2026-09-30 15:24:15 +08:00

405 lines
10 KiB
Go

package broker
import (
"context"
"database/sql"
"encoding/hex"
"errors"
"strings"
"sync"
"time"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
// ErrEndpointNotFound 编号不存在(业务拒绝,非内部故障)。
var ErrEndpointNotFound = errors.New("broker: endpoint not found")
// ErrEndpointDisabled 端已停用(业务拒绝)。
var ErrEndpointDisabled = errors.New("broker: endpoint disabled")
// Login 实现 Authenticator:会话令牌或登录密码(含锁定)。
type Login struct {
DB *store.DB
Pool auth.HashPool
Tokens auth.SessionTokens
Locks auth.LoginLocks
IdleDays int
Now func() time.Time
usedMu sync.Mutex
// 内存中的 session_used_at(毫秒)与上次落库时间。
usedAt map[string]int64
lastFlush map[string]int64
verifyMu sync.Mutex
verifySem map[string]chan struct{}
}
// LoginOptions 装配 Login。
type LoginOptions struct {
DB *store.DB
Pool auth.HashPool
Tokens auth.SessionTokens
Locks auth.LoginLocks
IdleDays int
Now func() time.Time
}
// NewLogin 创建登录校验器。
func NewLogin(opts LoginOptions) *Login {
now := opts.Now
if now == nil {
now = time.Now
}
tokens := opts.Tokens
if tokens == nil {
tokens = auth.NewSessionTokens()
}
locks := opts.Locks
if locks == nil {
locks = auth.NewLoginLocks()
}
return &Login{
DB: opts.DB,
Pool: opts.Pool,
Tokens: tokens,
Locks: locks,
IdleDays: opts.IdleDays,
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 {
return AuthResult{}, errors.New("broker: login not configured")
}
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 {
if errors.Is(err, ErrEndpointNotFound) || errors.Is(err, ErrEndpointDisabled) {
return AuthResult{OK: false}, nil
}
return AuthResult{}, err
}
cred := string(password)
if l.Tokens.LooksLikeSessionToken(cred) {
ok, authErr := l.authSession(ctx, endpointID, cred, row)
if authErr != nil {
return AuthResult{}, authErr
}
return AuthResult{OK: ok}, nil
}
ok, token, authErr := l.authPassword(ctx, endpointID, cred, remoteIP, row)
if authErr != nil {
return AuthResult{}, authErr
}
return AuthResult{OK: ok, SessionToken: token}, nil
}
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) {
var (
loginHash string
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, 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
}
return endpointAuthRow{}, err
}
if enabled == 0 {
return endpointAuthRow{}, ErrEndpointDisabled
}
row := endpointAuthRow{loginHash: loginHash}
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 {
// 损坏的哈希视为无有效会话(令牌校验失败),不是内部故障
row.sessionHash = nil
} else {
row.sessionHash = raw
}
}
return row, nil
}
func (l *Login) authSession(ctx context.Context, endpointID, token string, row endpointAuthRow) (bool, error) {
if len(row.sessionHash) == 0 {
return false, nil
}
got := l.Tokens.HashToken(token)
if !auth.EqualHash(got, row.sessionHash) {
return false, nil
}
now := l.Now()
nowMs := now.UnixMilli()
usedAt := row.sessionUsedAt
l.usedMu.Lock()
if mem, ok := l.usedAt[endpointID]; ok && mem > usedAt {
usedAt = mem
}
l.usedMu.Unlock()
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
}
return true, nil
}
func (l *Login) touchSessionUsed(ctx context.Context, endpointID string, nowMs int64) error {
const flushEvery = int64(time.Hour / time.Millisecond)
l.usedMu.Lock()
l.usedAt[endpointID] = nowMs
last := l.lastFlush[endpointID]
needFlush := last == 0 || nowMs-last >= flushEvery
if needFlush {
l.lastFlush[endpointID] = nowMs
}
l.usedMu.Unlock()
if !needFlush {
return nil
}
return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET session_used_at = ? WHERE id = ? AND session_hash IS NOT NULL AND session_hash != ''`,
nowMs, endpointID)
return err
})
}
func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP string, row endpointAuthRow) (ok bool, token string, err error) {
ipKey := auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}
epKey := auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}
if locked, _ := l.Locks.Check(ipKey); locked {
return false, "", nil
}
if locked, _ := l.Locks.Check(epKey); locked {
return false, "", nil
}
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
}
if !match {
l.Locks.Fail(ipKey)
l.Locks.Fail(epKey)
return false, "", nil
}
tok, hash, issErr := l.Tokens.Issue(ctx)
if issErr != nil {
return false, "", issErr
}
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 {
res, e := tx.Exec(`
UPDATE endpoints
SET session_hash = ?, session_issued_at = ?, session_used_at = ?
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
}
l.usedMu.Lock()
l.usedAt[endpointID] = nowMs
l.lastFlush[endpointID] = nowMs
l.usedMu.Unlock()
return true, tok, nil
}
// ClearSession 清空会话令牌(logout / 停用 / 删除 / 重置密码)。
func (l *Login) ClearSession(ctx context.Context, endpointID string) error {
if l == nil || l.DB == nil {
return errors.New("broker: login not configured")
}
err := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`
UPDATE endpoints
SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL
WHERE id = ?`, endpointID)
return e
})
if err != nil {
return err
}
l.usedMu.Lock()
delete(l.usedAt, endpointID)
delete(l.lastFlush, endpointID)
l.usedMu.Unlock()
return nil
}
// SetOnlineSince 握手完成时写入 online_since。
func (l *Login) SetOnlineSince(ctx context.Context, endpointID string, atMs int64) error {
return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET online_since = ? WHERE id = ?`, atMs, endpointID)
return err
})
}
// SetOfflineSince 当前连接断开时写入 offline_since。
func (l *Login) SetOfflineSince(ctx context.Context, endpointID string, atMs int64) error {
return l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id = ?`, atMs, endpointID)
return err
})
}
// SessionHashOf 返回当前库中的会话哈希(测试用);无则 nil。
func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, error) {
var sessHex sql.NullString
err := l.DB.Read.QueryRowContext(ctx, `SELECT session_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&sessHex)
if err != nil {
return nil, err
}
if !sessHex.Valid || sessHex.String == "" {
return nil, nil
}
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_")
}
var _ Authenticator = (*Login)(nil)