Files
NixMsg/internal/app/identity/talk.go
T

182 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
}
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)
})
}