feat: 实现身份资料、对话密码、在线目录与群管理

This commit is contained in:
Nixevol
2026-09-30 07:33:31 +08:00
parent bdd1d9e9f4
commit a78ab0d547
14 changed files with 2777 additions and 65 deletions
+131
View File
@@ -0,0 +1,131 @@
package identity
import (
"context"
"encoding/hex"
"log/slog"
"net/http"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
const (
grantKindPassword = "password"
grantKindReply = "reply"
)
// Config 是身份服务依赖(注册 + self + 对话密码)。
type Config struct {
DB *store.DB
Hash auth.HashPool
Locks auth.LoginLocks
Logger *slog.Logger
Now func() time.Time
// ClientIP 仅注册 HTTP 用。
ClientIP func(*http.Request) string
// Sessions 签发会话令牌;改登录密码必填。
Sessions auth.SessionTokens
// MaxScheduleSeconds 限制 self.update 的 default_delay_ms。
MaxScheduleSeconds int64
// ConnControl 可选:logout 后踢线;未接线时为 nil。
ConnControl port.ConnControl
}
// App 实现 identity.Service(含 I1 注册与 I2 self/对话密码)。
type App struct {
handler *RegisterHandler
db *store.DB
hash auth.HashPool
locks auth.LoginLocks
sessions auth.SessionTokens
maxScheduleSeconds int64
connCtrl port.ConnControl
nowFn func() time.Time
}
// New 构造完整身份服务。
func New(cfg Config) *App {
if cfg.Logger == nil {
cfg.Logger = slog.Default()
}
if cfg.Now == nil {
cfg.Now = time.Now
}
if cfg.ClientIP == nil {
cfg.ClientIP = clientIPFromRemoteAddr
}
if cfg.Sessions == nil {
cfg.Sessions = auth.NewStubSessionTokens()
}
if cfg.Locks == nil {
cfg.Locks = auth.NewStubLoginLocks()
}
h := NewRegisterHandler(RegisterConfig{
DB: cfg.DB,
Hash: cfg.Hash,
Locks: cfg.Locks,
Logger: cfg.Logger,
Now: cfg.Now,
ClientIP: cfg.ClientIP,
})
return &App{
handler: h,
db: cfg.DB,
hash: cfg.Hash,
locks: cfg.Locks,
sessions: cfg.Sessions,
maxScheduleSeconds: cfg.MaxScheduleSeconds,
connCtrl: cfg.ConnControl,
nowFn: cfg.Now,
}
}
// NewServer 兼容 I1:用注册配置构造 Service(会话令牌用 Stub)。
func NewServer(cfg RegisterConfig) *App {
return New(Config{
DB: cfg.DB,
Hash: cfg.Hash,
Locks: cfg.Locks,
Logger: cfg.Logger,
Now: cfg.Now,
ClientIP: cfg.ClientIP,
Sessions: auth.NewStubSessionTokens(),
})
}
func (a *App) now() time.Time { return a.nowFn() }
// Handler 返回可挂载的注册 HTTP 处理器。
func (a *App) Handler() http.Handler { return a.handler }
// Register 实现自助注册。
func (a *App) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
preq := &protocol.RegisterRequest{
RegistrationCode: req.RegistrationCode,
ID: req.ID,
LoginPassword: req.LoginPassword,
Name: req.Name,
TalkPassword: req.TalkPassword,
}
ip := req.RemoteIP
if ip == "" {
ip = "0.0.0.0"
}
result, apiErr := a.handler.register(ctx, preq, ip)
if apiErr != nil {
return RegisterResult{}, apiErr
}
return result, nil
}
func encodeSessionHash(hash []byte) string {
return hex.EncodeToString(hash)
}
var _ Service = (*App)(nil)
var _ http.Handler = (*RegisterHandler)(nil)
+7
View File
@@ -0,0 +1,7 @@
package identity
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
func errCode(code, msg string) *protocol.Error {
return &protocol.Error{Code: code, Message: msg}
}
-59
View File
@@ -247,65 +247,6 @@ func (h *RegisterHandler) logResult(result, id, ip string) {
h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip)
}
// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。
type Server struct {
handler *RegisterHandler
}
// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。
func NewServer(cfg RegisterConfig) *Server {
return &Server{handler: NewRegisterHandler(cfg)}
}
// Handler 返回可挂载的注册 HTTP 处理器。
func (s *Server) Handler() http.Handler { return s.handler }
// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。
func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
preq := &protocol.RegisterRequest{
RegistrationCode: req.RegistrationCode,
ID: req.ID,
LoginPassword: req.LoginPassword,
Name: req.Name,
TalkPassword: req.TalkPassword,
}
ip := req.RemoteIP
if ip == "" {
ip = "0.0.0.0"
}
result, err := s.handler.register(ctx, preq, ip)
if err != nil {
return RegisterResult{}, err
}
return result, nil
}
func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) {
return SelfInfo{}, ErrNotImplemented
}
func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error {
return ErrNotImplemented
}
func (s *Server) SelfSetTalkPassword(context.Context, string, string) error {
return ErrNotImplemented
}
func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
return "", ErrNotImplemented
}
func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented }
func (s *Server) UnlockTalk(context.Context, string, string, string) error {
return ErrNotImplemented
}
func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) {
return false, nil
}
func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented }
func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented }
func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented }
var _ Service = (*Server)(nil)
var _ http.Handler = (*RegisterHandler)(nil)
func setCORS(w http.ResponseWriter) {
w.Header().Set("Access-Control-Allow-Origin", "*")
}
+222
View File
@@ -0,0 +1,222 @@
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
}
// Disable / Enable / Delete 属 I5,此处保留未实现。
func (a *App) Disable(context.Context, string) error { return ErrNotImplemented }
func (a *App) Enable(context.Context, string) error { return ErrNotImplemented }
func (a *App) Delete(context.Context, string) error { return ErrNotImplemented }
+314
View File
@@ -0,0 +1,314 @@
package identity_test
import (
"context"
"database/sql"
"errors"
"path/filepath"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
"git.asio.asia/nixevol/NixMsg/internal/app/message"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
func openIdentity(t *testing.T) (*identity.App, *store.DB, *auth.MemoryLocks) {
t.Helper()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
locks := auth.NewLoginLocks()
locks.SetClock(func() time.Time { return fixed })
app := identity.New(identity.Config{
DB: db,
Hash: auth.NewStubHashPool(),
Locks: locks,
Sessions: auth.NewSessionTokens(),
MaxScheduleSeconds: int64(config.Default().Limits.MaxScheduleSeconds),
Now: func() time.Time { return fixed },
})
return app, db, locks
}
func insertEP(t *testing.T, db *store.DB, id, loginPW string) {
t.Helper()
err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, e := tx.Exec(`
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
VALUES(?,?,?,?,0,0,1,?)`, id, id, "stub$"+loginPW, nil, 1_700_000_000_000)
return e
})
if err != nil {
t.Fatal(err)
}
}
func protoCode(err error) string {
var pe *protocol.Error
if errors.As(err, &pe) {
return pe.Code
}
return ""
}
func TestF15TalkPasswordAuth(t *testing.T) {
t.Parallel()
app, db, _ := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
// A 不带密码失败
if err := app.UnlockTalk(ctx, "alice", "bob", "", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("want talk_password_required got %v", err)
}
ok, err := app.HasTalkGrant(ctx, "alice", "bob")
if err != nil || ok {
t.Fatalf("grant=%v err=%v", ok, err)
}
// 带对后成功,之后无密码也有授权
if err = app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
t.Fatal(err)
}
ok, err = app.HasTalkGrant(ctx, "alice", "bob")
if err != nil || !ok {
t.Fatalf("grant=%v err=%v", ok, err)
}
// B 改密后旧授权失效
if err = app.SelfSetTalkPassword(ctx, "bob", "newsecret"); err != nil {
t.Fatal(err)
}
ok, err = app.HasTalkGrant(ctx, "alice", "bob")
if err != nil || ok {
t.Fatalf("after change grant=%v err=%v", ok, err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("want invalid got %v", err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "newsecret", "1.1.1.1"); err != nil {
t.Fatal(err)
}
}
func TestF15ReplyGrantAndChange(t *testing.T) {
t.Parallel()
app, db, _ := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "alice", "alice-pw"); err != nil {
t.Fatal(err)
}
// B 先给 A 发 → 写入 reply 授权(A 可回 B 免密;这里记的是 bob→alice 的授权给「alice 作为接收方」...
// 回复授权:对方曾成功提交发给我的单聊 → 我对对方有 reply 权。
// 即 B 发给 A 后,A 对 B 有授权(sender=alice, target=bob)。
if err := app.RecordReplyGrant(ctx, "alice", "bob"); err != nil {
t.Fatal(err)
}
// 但 bob 还没设密码,alice→bob 本就不需要。给 bob 设密后验证 reply:
if err := app.SelfSetTalkPassword(ctx, "bob", "bob-pw"); err != nil {
t.Fatal(err)
}
// 重新记 reply:B 发给 A 成功后 A 获得对 B 的回复权
if err := app.RecordReplyGrant(ctx, "alice", "bob"); err != nil {
t.Fatal(err)
}
ok, err := app.HasTalkGrant(ctx, "alice", "bob")
if err != nil || !ok {
t.Fatalf("reply grant=%v err=%v", ok, err)
}
if err = app.SelfSetTalkPassword(ctx, "bob", "bob-pw2"); err != nil {
t.Fatal(err)
}
ok, err = app.HasTalkGrant(ctx, "alice", "bob")
if err != nil || ok {
t.Fatalf("after change reply should die grant=%v", ok)
}
}
func TestF15JoinNeedsPasswordDespiteGrant(t *testing.T) {
t.Parallel()
app, db, _ := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
t.Fatal(err)
}
// 已有单聊授权,进群仍要密码
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("want required got %v", err)
}
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "wrong", "1.1.1.1"); protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("want invalid got %v", err)
}
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "secret", "1.1.1.1"); err != nil {
t.Fatal(err)
}
}
func TestF15SubmittedMessageUnaffectedByPasswordChange(t *testing.T) {
t.Parallel()
idApp, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
lim := message.LimitsFromConfig(config.Default().Limits)
lim.RequestsPerSecond = 0
msgApp := message.New(db, lim, auth.NewStubHashPool(),
message.WithNow(func() time.Time { return time.UnixMilli(1_700_000_000_000) }),
message.WithLocks(locks),
)
req := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
DelayMs: ptrInt64(60_000),
TalkPassword: "secret",
}
res, err := msgApp.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "1.1.1.1"}, req)
if err != nil {
t.Fatal(err)
}
if res.State != message.StateScheduled {
t.Fatalf("state=%s", res.State)
}
if err := idApp.SelfSetTalkPassword(ctx, "bob", "changed"); err != nil {
t.Fatal(err)
}
// 已提交消息行仍在且状态不变
var state string
if err := db.Read.QueryRow(`SELECT state FROM messages WHERE sender_id=? AND id=?`, "alice", "m1").Scan(&state); err != nil {
t.Fatal(err)
}
if state != message.StateScheduled {
t.Fatalf("message state changed to %s", state)
}
}
func TestF15TargetLockAndExistingGrant(t *testing.T) {
t.Parallel()
app, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "bob", "password1")
insertEP(t, db, "authd", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
if err := app.UnlockTalk(ctx, "authd", "bob", "secret", "9.9.9.9"); err != nil {
t.Fatal(err)
}
// 50 次错误触发对方总数锁
for i := 0; i < 50; i++ {
id := "u" + string(rune('0'+i/100)) + string(rune('0'+(i/10)%10)) + string(rune('0'+i%10))
insertEP(t, db, id, "password1")
_ = app.UnlockTalk(ctx, id, "bob", "wrong", "2.2.2.2")
}
locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"})
if !locked {
t.Fatal("expected talk target lock")
}
// 正确密码也暂时无法解锁
insertEP(t, db, "newbie", "password1")
if err := app.UnlockTalk(ctx, "newbie", "bob", "secret", "3.3.3.3"); protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("want rate_limited got %v", err)
}
// 已有授权仍可用
ok, err := app.HasTalkGrant(ctx, "authd", "bob")
if err != nil || !ok {
t.Fatalf("authd grant=%v err=%v", ok, err)
}
// 改密清零
if err := app.SelfSetTalkPassword(ctx, "bob", "secret2"); err != nil {
t.Fatal(err)
}
locked, _ = locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"})
if locked {
t.Fatal("lock should clear on password change")
}
}
func TestSelfLoginPasswordAndLogout(t *testing.T) {
t.Parallel()
app, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "oldpass12")
tok, err := app.SelfChangeLoginPassword(ctx, "alice", "wrongpass", "newpass12", "1.1.1.1")
if protoCode(err) != protocol.CodeUnauthorized || tok != "" {
t.Fatalf("got tok=%q err=%v", tok, err)
}
locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: "alice", IP: "1.1.1.1"}) // 确保 Fail 路径可调用
tok, err = app.SelfChangeLoginPassword(ctx, "alice", "oldpass12", "newpass12", "1.1.1.1")
if err != nil || tok == "" || !protocol.ValidEndpointID("alice") {
t.Fatalf("tok=%q err=%v", tok, err)
}
if !auth.NewSessionTokens().LooksLikeSessionToken(tok) {
t.Fatalf("token prefix %q", tok)
}
var hash sql.NullString
if err = db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id=?`, "alice").Scan(&hash); err != nil || !hash.Valid {
t.Fatal(err)
}
info, err := app.SelfGet(ctx, "alice")
if err != nil || info.ID != "alice" {
t.Fatal(err)
}
name := "门口"
delay := int64(10000)
if err := app.SelfUpdate(ctx, "alice", &protocol.SelfUpdate{
V: protocol.Version, Type: protocol.TypeSelfUpdate, RID: "1", Name: name, DefaultDelayMs: &delay,
}); err != nil {
t.Fatal(err)
}
info, _ = app.SelfGet(ctx, "alice")
if info.Name != name || info.DefaultDelayMs != delay {
t.Fatalf("%+v", info)
}
if err := app.SelfLogout(ctx, "alice"); err != nil {
t.Fatal(err)
}
if err := db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id=?`, "alice").Scan(&hash); err != nil {
t.Fatal(err)
}
if hash.Valid {
t.Fatal("session should be cleared")
}
}
func TestUnlockSelfAndNoPassword(t *testing.T) {
t.Parallel()
app, db, _ := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.UnlockTalk(ctx, "alice", "alice", "", ""); err != nil {
t.Fatal(err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "anything", ""); err != nil {
t.Fatal(err)
}
}
func ptrInt64(v int64) *int64 { return &v }
+7 -4
View File
@@ -46,13 +46,16 @@ type Service interface {
SelfGet(ctx context.Context, endpointID string) (SelfInfo, error)
SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error
SelfSetTalkPassword(ctx context.Context, endpointID string, talkPassword string) error
SelfChangeLoginPassword(ctx context.Context, endpointID string, oldPassword, newPassword string) (sessionToken string, err error)
// SelfChangeLoginPassword 校验旧密码后换新密码并签发新会话令牌;remoteIP 计入登录锁定。
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
SelfLogout(ctx context.Context, endpointID string) error
// UnlockTalk 校验并写入对话密码授权(第 6.6 节 unlock)。
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword string) error
// HasTalkGrant 查询发送方对目标是否有有效授权。
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。
HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error)
// CheckTalkPasswordForJoin 加人时校验对话密码:已有单聊授权不能代替,必须当次带对。
CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error
// Disable 停用端并作废相关消息/令牌(第 7.6 节)。
Disable(ctx context.Context, endpointID string) error
+6 -2
View File
@@ -27,13 +27,13 @@ func (s *Stub) SelfSetTalkPassword(context.Context, string, string) error {
return ErrNotImplemented
}
func (s *Stub) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
func (s *Stub) SelfChangeLoginPassword(context.Context, string, string, string, string) (string, error) {
return "", ErrNotImplemented
}
func (s *Stub) SelfLogout(context.Context, string) error { return ErrNotImplemented }
func (s *Stub) UnlockTalk(context.Context, string, string, string) error {
func (s *Stub) UnlockTalk(context.Context, string, string, string, string) error {
return ErrNotImplemented
}
@@ -41,6 +41,10 @@ func (s *Stub) HasTalkGrant(context.Context, string, string) (bool, error) {
return false, nil
}
func (s *Stub) CheckTalkPasswordForJoin(context.Context, string, string, string, string) error {
return ErrNotImplemented
}
func (s *Stub) Disable(context.Context, string) error { return ErrNotImplemented }
func (s *Stub) Enable(context.Context, string) error { return ErrNotImplemented }
func (s *Stub) Delete(context.Context, string) error { return ErrNotImplemented }
+181
View File
@@ -0,0 +1,181 @@
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)
})
}