Files
NixMsg/internal/app/message/submit_test.go
T

387 lines
11 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 // 测试默认不限速
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("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)
}
})
}