package identity_test import ( "context" "database/sql" "encoding/json" "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 TestDisableEmitsRevokedForPushed(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) down := &message.RecordingDownlink{} ctrl := &port.StubConnControl{} 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: ctrl, Downlink: down, }) ctx := context.Background() insertEPFull(t, db, "alice") insertEPFull(t, db, "bob") err = db.Queue.Do(ctx, func(tx *sql.Tx) error { res, e := tx.Exec(` INSERT INTO messages( id, sender_id, dest_kind, dest_id, meta, content_type, body_enc, send_at, keep, ttl_seconds, receipt, state, reason, created_at) VALUES('pushed-1','alice','endpoint','bob','{}','text/plain','utf8',?,1,0,0,'dispatched','',?)`, fixed.UnixMilli(), fixed.UnixMilli()) if e != nil { return e } seq, _ := res.LastInsertId() _, e = tx.Exec(` INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at, pushed_at, pushed_conn) VALUES(?,?,?,1,'pending','',?,?,?)`, seq, "bob", fixed.UnixMilli(), fixed.UnixMilli(), fixed.UnixMilli(), "c-bob") return e }) if err != nil { t.Fatal(err) } if err := idApp.Disable(ctx, "bob"); err != nil { t.Fatal(err) } if down.FilterType(protocol.TypeRevoked) != 1 { t.Fatalf("want 1 revoked, got snapshots=%v", down.Snapshots()) } p := down.Snapshots()[0] var head struct { Type string `json:"type"` Reason string `json:"reason"` ID string `json:"id"` } _ = json.Unmarshal(p.Payload, &head) if head.Type != protocol.TypeRevoked || head.Reason != "endpoint_disabled" || head.ID != "pushed-1" { t.Fatalf("revoked=%+v", head) } if p.EndpointID != "bob" || p.QoS != 1 { t.Fatalf("publish=%+v", p) } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if len(ctrl.Calls) == 1 && ctrl.Calls[0] == "bob" { return } time.Sleep(5 * time.Millisecond) } t.Fatalf("disconnect calls=%v", ctrl.Calls) } 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 TestDisableLastPendingFinalizesAndReceipt(t *testing.T) { t.Parallel() idApp, msgApp, db := openLifecycle(t) ctx := context.Background() insertEPFull(t, db, "alice") insertEPFull(t, db, "bob") ttl := int64(3600) if _, err := msgApp.Submit(ctx, "alice", port.ConnInfo{EndpointID: "alice"}, &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "keep-bob", To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "bob"}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"}, Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl}, }); err != nil { t.Fatal(err) } if err := idApp.Disable(ctx, "bob"); err != nil { t.Fatal(err) } var state, reason string if err := db.Read.QueryRow(`SELECT state, reason FROM messages WHERE id='keep-bob'`).Scan(&state, &reason); err != nil { t.Fatal(err) } if state != "completed" { t.Fatalf("state=%s want completed", state) } var bodies int if err := db.Read.QueryRow(`SELECT COUNT(*) FROM message_bodies b JOIN messages m ON m.seq=b.seq WHERE m.id='keep-bob'`).Scan(&bodies); err != nil { t.Fatal(err) } if bodies != 0 { t.Fatalf("body still present: %d", bodies) } var rState, rReason string if err := db.Read.QueryRow(`SELECT state, reason FROM receipts WHERE sender_id='alice' AND msg_id='keep-bob'`).Scan(&rState, &rReason); err != nil { t.Fatalf("receipt: %v", err) } if rState != "rejected" || rReason != "endpoint_disabled" { t.Fatalf("receipt state=%q reason=%q", rState, rReason) } var pending int if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id='alice' AND state IN ('scheduled','dispatched')`).Scan(&pending); err != nil { t.Fatal(err) } if pending != 0 { t.Fatalf("sender pending=%d", pending) } } 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) } } func TestU03AdminTalkPasswordHTTP(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) } locks := auth.NewLoginLocks() locks.SetClock(func() time.Time { return fixed }) idApp := identity.New(identity.Config{ DB: db, Hash: hash, Locks: locks, Sessions: auth.NewSessionTokens(), MaxScheduleSeconds: 86400, Now: func() time.Time { return fixed }, }) h := admin.New(admin.Deps{ DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(), Locks: locks, 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) } req, _ := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints", strings.NewReader(`{"id":"alice","name":"alice","login_password":"password12"}`)) 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) } raw, _ := io.ReadAll(res.Body) _ = res.Body.Close() if res.StatusCode != 200 { t.Fatalf("create: %d %s", res.StatusCode, raw) } for i := 0; i < 50; i++ { locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}) } if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}); !locked { t.Fatal("expected target lock") } req, _ = http.NewRequest(http.MethodPut, srv.URL+"/api/admin/endpoints/alice/talk-password", strings.NewReader(`{"talk_password":"secret"}`)) 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) } raw, _ = io.ReadAll(res.Body) _ = res.Body.Close() if res.StatusCode != 200 { t.Fatalf("talk-password: %d %s", res.StatusCode, raw) } if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}); locked { t.Fatal("identity SelfSetTalkPassword should clear LockTalkTarget") } req, _ = http.NewRequest(http.MethodPut, srv.URL+"/api/admin/endpoints/missing/talk-password", strings.NewReader(`{"talk_password":"secret"}`)) 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) } raw, _ = io.ReadAll(res.Body) _ = res.Body.Close() if res.StatusCode != http.StatusNotFound { t.Fatalf("missing endpoint want 404 got %d %s", res.StatusCode, raw) } } type lifecycleKick struct { Calls []string } func (k *lifecycleKick) Kick(_ context.Context, endpointID string) (bool, error) { k.Calls = append(k.Calls, endpointID) return true, nil }