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 } // 没带密码不计锁定、也不因已锁返回 rate_limited(DEVELOPMENT 第 5 节:带密码的发送/进群才限流)。 if talkPassword == "" { return errCode(protocol.CodeTalkPasswordRequired, "talk password required") } if a.talkRateLimited(senderID, targetID) { return errCode(protocol.CodeRateLimited, "talk password locked") } 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.talkFail(senderID, targetID) return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") } a.talkClearPair(senderID, targetID) _ = 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 talkPassword == "" { return errCode(protocol.CodeTalkPasswordRequired, "talk password required") } if a.talkRateLimited(actorID, targetID) { return errCode(protocol.CodeRateLimited, "talk password locked") } 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.talkFail(actorID, targetID) return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") } a.talkClearPair(actorID, targetID) // 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。 _ = talkVer _ = remoteIP return nil } func talkPairKey(senderID, targetID string) auth.LockKey { return auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID} } func talkTargetKey(targetID string) auth.LockKey { return auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID} } func (a *App) talkRateLimited(senderID, targetID string) bool { if a.locks == nil { return false } if locked, _ := a.locks.Check(talkPairKey(senderID, targetID)); locked { return true } locked, _ := a.locks.Check(talkTargetKey(targetID)) return locked } func (a *App) talkFail(senderID, targetID string) { if a.locks == nil { return } a.locks.Fail(talkPairKey(senderID, targetID)) a.locks.Fail(talkTargetKey(targetID)) } func (a *App) talkClearPair(senderID, targetID string) { if a.locks == nil { return } a.locks.Clear(talkPairKey(senderID, targetID)) } 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) }) }