package message import ( "context" "database/sql" "errors" "fmt" "math" "path/filepath" "testing" "time" "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 TestStubSubmitNotImplemented(t *testing.T) { s := NewStub() _, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{}) if !errors.Is(err, ErrNotImplemented) { t.Fatalf("got %v", err) } } func TestStubRecoverNoop(t *testing.T) { s := NewStub() if err := s.RecoverOnStart(context.Background()); err != nil { t.Fatal(err) } } func openTestApp(t *testing.T, lim Limits) (*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) app := New(db, lim, auth.NewStubHashPool(), WithNow(func() time.Time { return fixed }), WithLocks(auth.NewStubLoginLocks()), ) return app, db } func defaultTestLimits() Limits { cfg := config.Default().Limits lim := LimitsFromConfig(cfg) lim.RequestsPerSecond = 0 // 测试默认不限速 lim.RecordRetentionDays = 7 lim.ReceiptRetentionDays = 7 lim.IdempotencyHours = 24 return lim } func insertEndpoint(t *testing.T, db *store.DB, id string, talkPassword string, enabled int, defaultDelayMs int64) { t.Helper() ctx := context.Background() var talk any var talkVer int64 if talkPassword != "" { talk = "stub$" + talkPassword talkVer = 1 } nowMs := int64(1_700_000_000_000) err := db.Queue.Do(ctx, 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(?,?,?,?,?,?,?,?,?)`, id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, nowMs, nowMs) return e }) if err != nil { t.Fatal(err) } } func baseSend(id, to string) *protocol.Send { return &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: id, To: protocol.Target{Kind: protocol.TargetEndpoint, ID: to}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"}, } } func protoCode(err error) string { var pe *protocol.Error if errors.As(err, &pe) { return pe.Code } return "" } func TestSubmitTable(t *testing.T) { t.Parallel() t.Run("idempotent_hit", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ctx := context.Background() req := baseSend("msg-1", "bob") first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } if first.State != StateDispatched { t.Fatalf("state=%s", first.State) } second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } if second != first { t.Fatalf("want %+v got %+v", first, second) } var n int if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id=? AND id=?`, "alice", "msg-1").Scan(&n); err != nil { t.Fatal(err) } if n != 1 { t.Fatalf("messages=%d", n) } }) t.Run("conflict", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ctx := context.Background() req := baseSend("msg-2", "bob") if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil { t.Fatal(err) } other := baseSend("msg-2", "bob") other.Body.Data = "other" _, err := app.Submit(ctx, "alice", port.ConnInfo{}, other) if protoCode(err) != protocol.CodeConflict { t.Fatalf("want conflict got %v", err) } }) t.Run("quota_exceeded", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() lim.MaxPendingPerSender = 1 app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ctx := context.Background() delay := int64(60_000) req1 := baseSend("q1", "bob") req1.DelayMs = &delay if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req1); err != nil { t.Fatal(err) } req2 := baseSend("q2", "bob") req2.DelayMs = &delay _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req2) if protoCode(err) != protocol.CodeQuotaExceeded { t.Fatalf("want quota_exceeded got %v", err) } }) t.Run("auth_required_and_grant", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "alice-secret", 1, 0) insertEndpoint(t, db, "bob", "secret", 1, 0) ctx := context.Background() _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a1", "bob")) if protoCode(err) != protocol.CodeTalkPasswordRequired { t.Fatalf("want talk_password_required got %v", err) } bad := baseSend("a2", "bob") bad.TalkPassword = "wrong" _, err = app.Submit(ctx, "alice", port.ConnInfo{}, bad) if protoCode(err) != protocol.CodeTalkPasswordInvalid { t.Fatalf("want talk_password_invalid got %v", err) } okReq := baseSend("a3", "bob") okReq.TalkPassword = "secret" res, err := app.Submit(ctx, "alice", port.ConnInfo{}, okReq) if err != nil { t.Fatal(err) } if res.State != StateDispatched { t.Fatalf("state=%s", res.State) } // 已有授权后不带密码也可发 if _, submitErr := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a4", "bob")); submitErr != nil { t.Fatal(submitErr) } // 回复授权:bob→alice(因 alice 设了对话密码) var kind string err = db.Read.QueryRow(` SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice").Scan(&kind) if err != nil { t.Fatal(err) } if kind != GrantKindReply { t.Fatalf("reply grant kind=%s", kind) } }) t.Run("self_skip_talk_password", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "secret", 1, 0) ctx := context.Background() if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("self1", "alice")); err != nil { t.Fatal(err) } }) t.Run("delay_and_send_at_mutex", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ctx := context.Background() delay := int64(1000) sendAt := int64(1_700_000_001_000) req := baseSend("m-mutex", "bob") req.DelayMs = &delay req.SendAtMs = &sendAt _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if protoCode(err) != protocol.CodeBadRequest { t.Fatalf("want bad_request got %v", err) } }) t.Run("idempotent_before_disabled_check", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ctx := context.Background() req := baseSend("pre-disable", "bob") first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } err = db.Queue.Do(ctx, func(tx *sql.Tx) error { _, e := tx.Exec(`UPDATE endpoints SET enabled = 0 WHERE id = ?`, "bob") return e }) if err != nil { t.Fatal(err) } // 新消息应失败 _, err = app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("after-disable", "bob")) if protoCode(err) != protocol.CodeEndpointDisabled { t.Fatalf("want endpoint_disabled got %v", err) } // 原请求重试仍返回原结果 second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } if second != first { t.Fatalf("want %+v got %+v", first, second) } }) t.Run("scheduled_not_dispatched", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ctx := context.Background() delay := int64(10_000) req := baseSend("sched-1", "bob") req.DelayMs = &delay res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } if res.State != StateScheduled { t.Fatalf("state=%s", res.State) } var n int if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil { t.Fatal(err) } if n != 0 { t.Fatalf("deliveries=%d", n) } }) t.Run("group_dispatch_excludes_sender", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) insertEndpoint(t, db, "carol", "", 1, 0) ctx := context.Background() 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", "g", "alice", 1_700_000_000_000); e != nil { return e } for _, m := range []string{"alice", "bob", "carol"} { if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`, "g1", m, 1_700_000_000_000); e != nil { return e } } return nil }) if err != nil { t.Fatal(err) } req := &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "gmsg-1", To: protocol.Target{Kind: protocol.TargetGroup, ID: "g1"}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"}, } res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } if res.State != StateDispatched { t.Fatalf("state=%s", res.State) } rows, err := db.Read.Query(`SELECT endpoint_id FROM deliveries ORDER BY endpoint_id`) if err != nil { t.Fatal(err) } defer func() { _ = rows.Close() }() var got []string for rows.Next() { var id string if err := rows.Scan(&id); err != nil { t.Fatal(err) } got = append(got, id) } if len(got) != 2 || got[0] != "bob" || got[1] != "carol" { t.Fatalf("recipients=%v", got) } }) t.Run("ttl_zero_rejected", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ttl := int64(0) req := baseSend("ttl0", "bob") req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl} _, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req) if protoCode(err) != protocol.CodeBadRequest { t.Fatalf("ttl=0 want bad_request got %v", err) } }) t.Run("ttl_negative_rejected", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) ttl := int64(-1) req := baseSend("ttlneg", "bob") req.Offline = &protocol.OfflineOpts{Keep: true, TTLSeconds: &ttl} _, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req) if protoCode(err) != protocol.CodeBadRequest { t.Fatalf("ttl=-1 want bad_request got %v", err) } }) t.Run("delay_maxint64_rejected", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) delay := int64(math.MaxInt64) req := baseSend("delaymax", "bob") req.DelayMs = &delay _, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, req) if protoCode(err) != protocol.CodeBadRequest { t.Fatalf("delay=MaxInt64 want bad_request got %v", err) } }) t.Run("sender_disabled_unauthorized", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 0, 0) insertEndpoint(t, db, "bob", "", 1, 0) _, err := app.Submit(context.Background(), "alice", port.ConnInfo{}, baseSend("from-off", "bob")) if protoCode(err) != protocol.CodeUnauthorized { t.Fatalf("want unauthorized got %v", err) } var n int if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil { t.Fatal(err) } if n != 0 { t.Fatalf("deliveries=%d", n) } }) t.Run("group_late_joiner_skipped", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) insertEndpoint(t, db, "dave", "", 1, 0) ctx := context.Background() nowMs := int64(1_700_000_000_000) err := db.Queue.Do(ctx, func(tx *sql.Tx) error { if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`, "g-late", "g", "alice", nowMs); e != nil { return e } for _, m := range []string{"alice", "bob"} { if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`, "g-late", m, nowMs); e != nil { return e } } return nil }) if err != nil { t.Fatal(err) } delay := int64(10_000) req := &protocol.Send{ V: protocol.Version, Type: protocol.TypeSend, RID: "r1", ID: "late-1", To: protocol.Target{Kind: protocol.TargetGroup, ID: "g-late"}, Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"}, DelayMs: &delay, Offline: keepTrue(), } res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req) if err != nil { t.Fatal(err) } if res.State != StateScheduled { t.Fatalf("state=%s", res.State) } err = db.Queue.Do(ctx, func(tx *sql.Tx) error { _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`, "g-late", "dave", res.SendAtMs+500) return e }) if err != nil { t.Fatal(err) } if _, err := app.DispatchDue(ctx, res.SendAtMs, 10); err != nil { t.Fatal(err) } var daveN int if err := db.Read.QueryRow(` SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='late-1' AND d.endpoint_id='dave'`).Scan(&daveN); err != nil { t.Fatal(err) } if daveN != 0 { t.Fatalf("late joiner deliveries=%d", daveN) } var bobN int if err := db.Read.QueryRow(` SELECT COUNT(*) FROM deliveries d JOIN messages m ON m.seq=d.seq WHERE m.id='late-1' AND d.endpoint_id='bob'`).Scan(&bobN); err != nil { t.Fatal(err) } if bobN != 1 { t.Fatalf("bob deliveries=%d", bobN) } }) t.Run("rate_limited", func(t *testing.T) { t.Parallel() lim := defaultTestLimits() lim.RequestsPerSecond = 50 lim.RequestBurst = 2 app, db := openTestApp(t, lim) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "", 1, 0) if !app.AllowRequest("alice") || !app.AllowRequest("alice") { t.Fatal("burst should allow first two") } if app.AllowRequest("alice") { t.Fatal("third request should be rate limited") } ctx := context.Background() if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil { t.Fatal(err) } }) } func TestU03SubmitTalkLockNoIP(t *testing.T) { t.Parallel() lim := defaultTestLimits() 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 := New(db, lim, auth.NewStubHashPool(), WithNow(func() time.Time { return fixed }), WithLocks(locks), ) insertEndpoint(t, db, "alice", "", 1, 0) insertEndpoint(t, db, "bob", "secret", 1, 0) ctx := context.Background() for i := 0; i < 10; i++ { bad := baseSend(fmt.Sprintf("w%d", i), "bob") bad.TalkPassword = "wrong" _, err := app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: fmt.Sprintf("10.0.0.%d", i+1)}, bad) if protoCode(err) != protocol.CodeTalkPasswordInvalid { t.Fatalf("i=%d got %v", i, err) } } empty := baseSend("empty", "bob") _, err = app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "8.8.8.8"}, empty) if protoCode(err) != protocol.CodeTalkPasswordRequired { t.Fatalf("empty while locked want required got %v", err) } okReq := baseSend("ok1", "bob") okReq.TalkPassword = "secret" _, err = app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "9.9.9.9"}, okReq) if protoCode(err) != protocol.CodeRateLimited { t.Fatalf("correct password while locked want rate_limited got %v", err) } }