feat: 实现登录、会话令牌、握手与顶号
EOF
This commit is contained in:
@@ -0,0 +1,296 @@
|
||||
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
|
||||
}
|
||||
|
||||
// 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),
|
||||
}
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) {
|
||||
var (
|
||||
loginHash string
|
||||
enabled int
|
||||
sessHex sql.NullString
|
||||
usedAt 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)
|
||||
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 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 {
|
||||
idle := time.Duration(l.IdleDays) * 24 * time.Hour
|
||||
if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle {
|
||||
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")
|
||||
}
|
||||
match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash)
|
||||
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)
|
||||
writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`
|
||||
UPDATE endpoints
|
||||
SET session_hash = ?, session_issued_at = ?, session_used_at = ?
|
||||
WHERE id = ?`, hashHex, nowMs, nowMs, endpointID)
|
||||
return e
|
||||
})
|
||||
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)
|
||||
}
|
||||
|
||||
// LooksLikeSessionToken 暴露给测试。
|
||||
func (l *Login) LooksLikeSessionToken(s string) bool {
|
||||
return strings.HasPrefix(s, "nst_")
|
||||
}
|
||||
|
||||
var _ Authenticator = (*Login)(nil)
|
||||
Reference in New Issue
Block a user