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