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)