package identity import ( "context" "database/sql" "errors" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) // UnlockTalk 校验对话密码并写入 password 授权;未设密码则直接成功。 func (a *App) UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error { if !protocol.ValidEndpointID(senderID) || !protocol.ValidEndpointID(targetID) { return errCode(protocol.CodeBadRequest, "invalid endpoint id") } if senderID == targetID { return nil } var talkHash sql.NullString var talkVer int64 err := a.db.Read.QueryRowContext(ctx, ` SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer) if errors.Is(err, sql.ErrNoRows) { return errCode(protocol.CodeInvalidTarget, "target not found") } if err != nil { return err } if !talkHash.Valid || talkHash.String == "" { return nil } if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}); locked { return errCode(protocol.CodeRateLimited, "talk password locked") } if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked { return errCode(protocol.CodeRateLimited, "talk password locked") } if talkPassword == "" { return errCode(protocol.CodeTalkPasswordRequired, "talk password required") } if a.hash == nil { return errors.New("identity: hash pool required") } ok, err := a.hash.Verify(ctx, auth.PasswordTalk, talkPassword, talkHash.String) if err != nil { return err } if !ok { a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}) return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") } a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) nowMs := a.now().UnixMilli() return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { return upsertGrantTx(tx, senderID, targetID, talkVer, grantKindPassword, nowMs) }) } // HasTalkGrant 发给自己、对方未设密码、或存在匹配版本授权时为 true。 func (a *App) HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) { if senderID == targetID { return true, nil } if !protocol.ValidEndpointID(senderID) || !protocol.ValidEndpointID(targetID) { return false, errCode(protocol.CodeBadRequest, "invalid endpoint id") } var talkHash sql.NullString var talkVer int64 err := a.db.Read.QueryRowContext(ctx, ` SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer) if errors.Is(err, sql.ErrNoRows) { return false, errCode(protocol.CodeInvalidTarget, "target not found") } if err != nil { return false, err } if !talkHash.Valid || talkHash.String == "" { return true, nil } var n int err = a.db.Read.QueryRowContext(ctx, ` SELECT 1 FROM talk_grants WHERE sender_id = ? AND target_id = ? AND target_talk_version = ? LIMIT 1`, senderID, targetID, talkVer).Scan(&n) if errors.Is(err, sql.ErrNoRows) { return false, nil } return err == nil, err } // CheckTalkPasswordForJoin 加人时必须当次带对密码;已有单聊授权不能代替。 func (a *App) CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error { if !protocol.ValidEndpointID(targetID) { return errCode(protocol.CodeInvalidTarget, "invalid target") } var talkHash sql.NullString var talkVer int64 var enabled int err := a.db.Read.QueryRowContext(ctx, ` SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer, &enabled) if errors.Is(err, sql.ErrNoRows) { return errCode(protocol.CodeInvalidTarget, "target not found") } if err != nil { return err } if enabled == 0 { return errCode(protocol.CodeEndpointDisabled, "endpoint disabled") } if !talkHash.Valid || talkHash.String == "" { return nil } if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}); locked { return errCode(protocol.CodeRateLimited, "talk password locked") } if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked { return errCode(protocol.CodeRateLimited, "talk password locked") } if talkPassword == "" { return errCode(protocol.CodeTalkPasswordRequired, "talk password required") } if a.hash == nil { return errors.New("identity: hash pool required") } ok, err := a.hash.Verify(ctx, auth.PasswordTalk, talkPassword, talkHash.String) if err != nil { return err } if !ok { a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}) return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") } a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) // 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。 _ = talkVer return nil } func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error { _, err := tx.Exec(` INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at) VALUES(?,?,?,?,?) ON CONFLICT(sender_id, target_id) DO UPDATE SET target_talk_version = excluded.target_talk_version, kind = excluded.kind, created_at = excluded.created_at`, senderID, targetID, talkVersion, kind, nowMs, ) return err } // RecordReplyGrant 在对方成功提交单聊后写入 reply 授权(供消息线或测试调用)。 func (a *App) RecordReplyGrant(ctx context.Context, senderID, targetID string) error { if senderID == targetID { return nil } var talkHash sql.NullString var talkVer int64 err := a.db.Read.QueryRowContext(ctx, ` SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&talkHash, &talkVer) if errors.Is(err, sql.ErrNoRows) { return errCode(protocol.CodeInvalidTarget, "target not found") } if err != nil { return err } if !talkHash.Valid || talkHash.String == "" { return nil } nowMs := a.now().UnixMilli() return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { return upsertGrantTx(tx, senderID, targetID, talkVer, grantKindReply, nowMs) }) }