383 lines
10 KiB
Go
383 lines
10 KiB
Go
package message
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"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 // 测试默认不限速
|
|
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
|
|
}
|
|
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)
|
|
VALUES(?,?,?,?,?,?,?,?)`,
|
|
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, 1_700_000_000_000)
|
|
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("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)
|
|
ctx := context.Background()
|
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r2", "bob")); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r3", "bob"))
|
|
if protoCode(err) != protocol.CodeRateLimited {
|
|
t.Fatalf("want rate_limited got %v", err)
|
|
}
|
|
})
|
|
}
|