211 lines
6.4 KiB
Go
211 lines
6.4 KiB
Go
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)
|
||
})
|
||
}
|