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

387 lines
12 KiB
Go

package identity_test
import (
"context"
"database/sql"
"io"
"net/http"
"net/http/cookiejar"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/admin"
"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 openLifecycle(t *testing.T) (*identity.App, *message.App, *store.DB) {
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)
idApp := identity.New(identity.Config{
DB: db,
Hash: auth.NewStubHashPool(),
Locks: auth.NewStubLoginLocks(),
Sessions: auth.NewSessionTokens(),
MaxScheduleSeconds: int64(config.Default().Limits.MaxScheduleSeconds),
Now: func() time.Time { return fixed },
ConnControl: &port.StubConnControl{},
Downlink: &port.StubDownlink{},
})
lim := message.LimitsFromConfig(config.Default().Limits)
lim.RequestsPerSecond = 0
lim.RecordRetentionDays = 7
msgApp := message.New(db, lim, auth.NewStubHashPool(),
message.WithNow(func() time.Time { return fixed }),
message.WithLocks(auth.NewStubLoginLocks()),
)
return idApp, msgApp, db
}
func insertEPFull(t *testing.T, db *store.DB, id 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, offline_since)
VALUES(?,?,?,?,0,0,1,?,?)`, id, id, "stub$login", nil, 1_700_000_000_000, 1_700_000_000_000)
return e
})
if err != nil {
t.Fatal(err)
}
}
func TestF01DisableVoidsScheduledAndRejectsNew(t *testing.T) {
t.Parallel()
idApp, msgApp, db := openLifecycle(t)
ctx := context.Background()
insertEPFull(t, db, "alice")
insertEPFull(t, db, "bob")
future := int64(1_700_000_000_000 + 3600_000)
receipt := true
toBob := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m-to-bob",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
SendAtMs: &future,
Receipt: &receipt,
}
if _, err := msgApp.Submit(ctx, "alice", port.ConnInfo{EndpointID: "alice"}, toBob); err != nil {
t.Fatal(err)
}
fromBob := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "2", ID: "m-from-bob",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "alice"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "bye"},
SendAtMs: &future,
}
if _, err := msgApp.Submit(ctx, "bob", port.ConnInfo{EndpointID: "bob"}, fromBob); err != nil {
t.Fatal(err)
}
if err := idApp.Disable(ctx, "bob"); err != nil {
t.Fatal(err)
}
var enabled int
if err := db.Read.QueryRow(`SELECT enabled FROM endpoints WHERE id='bob'`).Scan(&enabled); err != nil {
t.Fatal(err)
}
if enabled != 0 {
t.Fatalf("enabled=%d", enabled)
}
var toState, toReason string
if err := db.Read.QueryRow(`SELECT state, reason FROM messages WHERE sender_id='alice' AND id='m-to-bob'`).
Scan(&toState, &toReason); err != nil {
t.Fatal(err)
}
if toState != "completed" || toReason != "endpoint_disabled" {
t.Fatalf("to bob: state=%s reason=%s", toState, toReason)
}
var fromState, fromReason string
if err := db.Read.QueryRow(`SELECT state, reason FROM messages WHERE sender_id='bob' AND id='m-from-bob'`).
Scan(&fromState, &fromReason); err != nil {
t.Fatal(err)
}
if fromState != "completed" || fromReason != "sender_disabled" {
t.Fatalf("from bob: state=%s reason=%s", fromState, fromReason)
}
newSend := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "3", ID: "m-new",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
}
_, err := msgApp.Submit(ctx, "alice", port.ConnInfo{EndpointID: "alice"}, newSend)
if protoCode(err) != protocol.CodeEndpointDisabled {
t.Fatalf("want endpoint_disabled got %v", err)
}
if err := idApp.Enable(ctx, "bob"); err != nil {
t.Fatal(err)
}
if err := db.Read.QueryRow(`SELECT state FROM messages WHERE id='m-to-bob'`).Scan(&toState); err != nil {
t.Fatal(err)
}
if toState != "completed" {
t.Fatalf("voided message restored? %s", toState)
}
okSend := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "4", ID: "m-after",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "ok"},
}
if _, err := msgApp.Submit(ctx, "alice", port.ConnInfo{EndpointID: "alice"}, okSend); err != nil {
t.Fatal(err)
}
}
func TestF01DeleteOwnerTransfersEarliest(t *testing.T) {
t.Parallel()
idApp, _, db := openLifecycle(t)
ctx := context.Background()
insertEPFull(t, db, "owner")
insertEPFull(t, db, "early")
insertEPFull(t, db, "late")
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES('g1','群','owner',?)`, 1_700_000_000_000); e != nil {
return e
}
_, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES
('g1','owner',100),('g1','early',200),('g1','late',300)`)
return e
})
if err != nil {
t.Fatal(err)
}
if err := idApp.Delete(ctx, "owner"); err != nil {
t.Fatal(err)
}
var owner string
if err := db.Read.QueryRow(`SELECT owner_id FROM groups WHERE id='g1'`).Scan(&owner); err != nil {
t.Fatal(err)
}
if owner != "early" {
t.Fatalf("want earliest other member early, got %s", owner)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id='g1' AND endpoint_id='owner'`).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatal("owner should have left the group")
}
}
func TestF01DeleteReopenNoOldReceipts(t *testing.T) {
t.Parallel()
idApp, msgApp, db := openLifecycle(t)
ctx := context.Background()
insertEPFull(t, db, "alice")
insertEPFull(t, db, "bob")
receipt := true
send := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "alice"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
Receipt: &receipt,
}
if _, err := msgApp.Submit(ctx, "bob", port.ConnInfo{EndpointID: "bob"}, send); err != nil {
t.Fatal(err)
}
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
VALUES('bob','m1','alice','accepted','',?,0)`, 1_700_000_000_000)
return e
})
if err != nil {
t.Fatal(err)
}
if err := idApp.Delete(ctx, "bob"); err != nil {
t.Fatal(err)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM receipts WHERE sender_id='bob'`).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("old receipts should be gone, got %d", n)
}
insertEPFull(t, db, "bob")
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM receipts WHERE sender_id='bob'`).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("reopened endpoint must not inherit receipts, got %d", n)
}
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id='bob'`).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("reopened endpoint must not inherit messages, got %d", n)
}
}
func TestAdminDisableDeleteHTTP(t *testing.T) {
t.Parallel()
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)
hash := auth.NewStubHashPool()
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, "adminpassword1"); seedErr != nil {
t.Fatal(seedErr)
}
idApp := identity.New(identity.Config{
DB: db, Hash: hash, Locks: auth.NewStubLoginLocks(), Sessions: auth.NewSessionTokens(),
MaxScheduleSeconds: 86400, Now: func() time.Time { return fixed },
})
lim := message.LimitsFromConfig(config.Default().Limits)
lim.RequestsPerSecond = 0
msgApp := message.New(db, lim, hash,
message.WithNow(func() time.Time { return fixed }),
message.WithLocks(auth.NewStubLoginLocks()),
)
kick := &lifecycleKick{}
h := admin.New(admin.Deps{
DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(),
Locks: admin.NewMemoryLoginLocks(), KickEndpoint: kick.Kick, Identity: idApp,
})
srv := httptest.NewServer(h)
t.Cleanup(srv.Close)
jar, _ := cookiejar.New(nil)
client := &http.Client{Jar: jar}
loginRes, err := client.Post(srv.URL+"/api/admin/login", "application/json",
strings.NewReader(`{"username":"admin","password":"adminpassword1"}`))
if err != nil {
t.Fatal(err)
}
_ = loginRes.Body.Close()
if loginRes.StatusCode != 200 {
t.Fatalf("login %d", loginRes.StatusCode)
}
createEP := func(id string) {
t.Helper()
req, _ := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints",
strings.NewReader(`{"id":"`+id+`","name":"`+id+`","login_password":"password12"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Nixmsg-Request", "1")
res, e := client.Do(req)
if e != nil {
t.Fatal(e)
}
raw, _ := io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != 200 {
t.Fatalf("create %s: %d %s", id, res.StatusCode, raw)
}
}
createEP("alice")
createEP("bob")
createEP("carol")
ctx := context.Background()
future := fixed.UnixMilli() + 3600_000
toBob := &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "sched-bob",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
SendAtMs: &future,
}
if _, subErr := msgApp.Submit(ctx, "alice", port.ConnInfo{EndpointID: "alice"}, toBob); subErr != nil {
t.Fatal(subErr)
}
req, _ := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints/batch",
strings.NewReader(`{"ids":["bob"],"action":"disable"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Nixmsg-Request", "1")
res, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != 200 {
t.Fatalf("disable: %d %s", res.StatusCode, body)
}
var st string
if scanErr := db.Read.QueryRow(`SELECT state FROM messages WHERE id='sched-bob'`).Scan(&st); scanErr != nil {
t.Fatal(scanErr)
}
if st != "completed" {
t.Fatalf("scheduled should be voided, got %s", st)
}
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES('ghttp','G','bob',?)`, fixed.UnixMilli()); e != nil {
return e
}
_, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES
('ghttp','bob',1),('ghttp','alice',2),('ghttp','carol',3)`)
return e
})
if err != nil {
t.Fatal(err)
}
// bob 已停用,需先启用才能作为「仍存在的群主」再删除?删除不要求 enabled。
req, _ = http.NewRequest(http.MethodDelete, srv.URL+"/api/admin/endpoints/bob", nil)
req.Header.Set("X-Nixmsg-Request", "1")
res, err = client.Do(req)
if err != nil {
t.Fatal(err)
}
raw, _ := io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != 200 {
t.Fatalf("delete: %d %s", res.StatusCode, raw)
}
var owner string
if err := db.Read.QueryRow(`SELECT owner_id FROM groups WHERE id='ghttp'`).Scan(&owner); err != nil {
t.Fatal(err)
}
if owner != "alice" {
t.Fatalf("want alice as new owner, got %s", owner)
}
}
type lifecycleKick struct {
Calls []string
}
func (k *lifecycleKick) Kick(_ context.Context, endpointID string) (bool, error) {
k.Calls = append(k.Calls, endpointID)
return true, nil
}