218 lines
6.7 KiB
Go
218 lines
6.7 KiB
Go
package identity
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
|
)
|
|
|
|
// SelfGet 返回自己的资料。
|
|
func (a *App) SelfGet(ctx context.Context, endpointID string) (SelfInfo, error) {
|
|
if !protocol.ValidEndpointID(endpointID) {
|
|
return SelfInfo{}, errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
|
}
|
|
var info SelfInfo
|
|
var talk sql.NullString
|
|
err := a.db.Read.QueryRowContext(ctx, `
|
|
SELECT id, name, default_delay_ms, talk_hash
|
|
FROM endpoints WHERE id = ?`, endpointID).Scan(&info.ID, &info.Name, &info.DefaultDelayMs, &talk)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return SelfInfo{}, errCode(protocol.CodeNotFound, "endpoint not found")
|
|
}
|
|
if err != nil {
|
|
return SelfInfo{}, err
|
|
}
|
|
info.TalkPasswordSet = talk.Valid && talk.String != ""
|
|
return info, nil
|
|
}
|
|
|
|
// SelfUpdate 更新名称与默认延迟。
|
|
func (a *App) SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error {
|
|
if !protocol.ValidEndpointID(endpointID) {
|
|
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
|
}
|
|
if req == nil {
|
|
return errCode(protocol.CodeBadRequest, "nil request")
|
|
}
|
|
if err := req.Validate(); err != nil {
|
|
return err
|
|
}
|
|
if req.Name == "" && req.DefaultDelayMs == nil {
|
|
return errCode(protocol.CodeBadRequest, "nothing to update")
|
|
}
|
|
if req.DefaultDelayMs != nil && a.maxScheduleSeconds > 0 {
|
|
maxMs := a.maxScheduleSeconds * 1000
|
|
if *req.DefaultDelayMs > maxMs {
|
|
return errCode(protocol.CodeBadRequest, "default_delay_ms exceeds max_schedule_seconds")
|
|
}
|
|
}
|
|
|
|
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
var exists int
|
|
if err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, endpointID).Scan(&exists); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return errCode(protocol.CodeNotFound, "endpoint not found")
|
|
}
|
|
return err
|
|
}
|
|
if req.Name != "" {
|
|
if _, err := tx.Exec(`UPDATE endpoints SET name = ? WHERE id = ?`, req.Name, endpointID); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
if req.DefaultDelayMs != nil {
|
|
if _, err := tx.Exec(`UPDATE endpoints SET default_delay_ms = ? WHERE id = ?`, *req.DefaultDelayMs, endpointID); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
})
|
|
}
|
|
|
|
// SelfSetTalkPassword 设置或清除对话密码,并增加 talk_version;改密清零对方锁定计数。
|
|
func (a *App) SelfSetTalkPassword(ctx context.Context, endpointID, talkPassword string) error {
|
|
if !protocol.ValidEndpointID(endpointID) {
|
|
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
|
}
|
|
if !protocol.ValidTalkPassword(talkPassword) {
|
|
return errCode(protocol.CodeBadRequest, "invalid talk_password")
|
|
}
|
|
|
|
var talkHash any
|
|
if talkPassword != "" {
|
|
if a.hash == nil {
|
|
return errors.New("identity: hash pool required")
|
|
}
|
|
phc, err := a.hash.Hash(ctx, auth.PasswordTalk, talkPassword)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
talkHash = phc
|
|
}
|
|
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
res, err := tx.Exec(`
|
|
UPDATE endpoints
|
|
SET talk_hash = ?, talk_version = talk_version + 1
|
|
WHERE id = ?`, talkHash, endpointID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
n, _ := res.RowsAffected()
|
|
if n == 0 {
|
|
return errCode(protocol.CodeNotFound, "endpoint not found")
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
// D24:改密清零按对方计的对话密码失败计数。
|
|
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: endpointID})
|
|
return nil
|
|
}
|
|
|
|
// SelfChangeLoginPassword 要求旧密码;成功返回新 session_token,旧令牌作废,当前连接由调用方保留。
|
|
func (a *App) SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (string, error) {
|
|
if !protocol.ValidEndpointID(endpointID) {
|
|
return "", errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
|
}
|
|
if protocol.LoginPasswordForbiddenPrefix(newPassword) {
|
|
return "", errCode(protocol.CodeBadRequest, "login password must not start with nst_")
|
|
}
|
|
if !protocol.ValidLoginPassword(newPassword) || newPassword == "" {
|
|
return "", errCode(protocol.CodeBadRequest, "invalid new_password")
|
|
}
|
|
if a.hash == nil || a.sessions == nil {
|
|
return "", errors.New("identity: hash/sessions required")
|
|
}
|
|
|
|
var loginHash string
|
|
err := a.db.Read.QueryRowContext(ctx, `SELECT login_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&loginHash)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return "", errCode(protocol.CodeNotFound, "endpoint not found")
|
|
}
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}); locked {
|
|
return "", errCode(protocol.CodeRateLimited, "login locked")
|
|
}
|
|
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}); locked {
|
|
return "", errCode(protocol.CodeRateLimited, "login locked")
|
|
}
|
|
|
|
ok, err := a.hash.Verify(ctx, auth.PasswordLogin, oldPassword, loginHash)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if !ok {
|
|
a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP})
|
|
a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID})
|
|
return "", errCode(protocol.CodeUnauthorized, "old password invalid")
|
|
}
|
|
|
|
newHash, err := a.hash.Hash(ctx, auth.PasswordLogin, newPassword)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
token, tokenHash, err := a.sessions.Issue(ctx)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
nowMs := a.now().UnixMilli()
|
|
hashHex := encodeSessionHash(tokenHash)
|
|
|
|
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
res, e := tx.Exec(`
|
|
UPDATE endpoints
|
|
SET login_hash = ?, session_hash = ?, session_issued_at = ?, session_used_at = ?
|
|
WHERE id = ?`, newHash, hashHex, nowMs, nowMs, endpointID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
n, _ := res.RowsAffected()
|
|
if n == 0 {
|
|
return errCode(protocol.CodeNotFound, "endpoint not found")
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return token, nil
|
|
}
|
|
|
|
// SelfLogout 清空会话令牌;若注入了 ConnControl 则断开当前连接。
|
|
func (a *App) SelfLogout(ctx context.Context, endpointID string) error {
|
|
if !protocol.ValidEndpointID(endpointID) {
|
|
return errCode(protocol.CodeBadRequest, "invalid endpoint id")
|
|
}
|
|
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
res, e := tx.Exec(`
|
|
UPDATE endpoints
|
|
SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL
|
|
WHERE id = ?`, endpointID)
|
|
if e != nil {
|
|
return e
|
|
}
|
|
n, _ := res.RowsAffected()
|
|
if n == 0 {
|
|
return errCode(protocol.CodeNotFound, "endpoint not found")
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if a.connCtrl != nil {
|
|
_ = a.connCtrl.Disconnect(ctx, endpointID, "", port.DisconnectNormal)
|
|
}
|
|
return nil
|
|
}
|