@@ -1,131 +0,0 @@
|
||||
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)
|
||||
@@ -1,7 +0,0 @@
|
||||
package identity
|
||||
|
||||
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
|
||||
func errCode(code, msg string) *protocol.Error {
|
||||
return &protocol.Error{Code: code, Message: msg}
|
||||
}
|
||||
@@ -247,6 +247,65 @@ 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", "*")
|
||||
}
|
||||
|
||||
@@ -1,222 +0,0 @@
|
||||
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 }
|
||||
@@ -1,314 +0,0 @@
|
||||
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 }
|
||||
@@ -46,16 +46,13 @@ 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 校验旧密码后换新密码并签发新会话令牌;remoteIP 计入登录锁定。
|
||||
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
|
||||
SelfChangeLoginPassword(ctx context.Context, endpointID string, oldPassword, newPassword string) (sessionToken string, err error)
|
||||
SelfLogout(ctx context.Context, endpointID string) error
|
||||
|
||||
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。
|
||||
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error
|
||||
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。
|
||||
// UnlockTalk 校验并写入对话密码授权(第 6.6 节 unlock)。
|
||||
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword 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
|
||||
|
||||
@@ -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) (string, error) {
|
||||
func (s *Stub) SelfChangeLoginPassword(context.Context, 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, string) error {
|
||||
func (s *Stub) UnlockTalk(context.Context, string, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
|
||||
@@ -41,10 +41,6 @@ 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 }
|
||||
|
||||
@@ -1,181 +0,0 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user