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 }