Files

536 lines
18 KiB
Go

package protocol_test
import (
"bytes"
"encoding/json"
"strings"
"testing"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
func ptrInt(v int) *int { return &v }
func ptrInt64(v int64) *int64 { return &v }
func ptrBool(v bool) *bool { return &v }
func mustMarshal(t *testing.T, v any) []byte {
t.Helper()
b, err := protocol.Marshal(v)
if err != nil {
t.Fatalf("Marshal: %v", err)
}
return b
}
func TestEncodeNoHTMLEscapeAndNoNewline(t *testing.T) {
msg := &protocol.Msg{
V: protocol.Version,
Type: protocol.TypeMsg,
ID: "id-1",
From: "app-1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "你好<script>&"},
}
b := mustMarshal(t, msg)
if bytes.Contains(b, []byte(`\u`)) {
t.Fatalf("unexpected unicode escape: %s", b)
}
if bytes.Contains(b, []byte(`\u003c`)) || !bytes.Contains(b, []byte("<script>")) {
t.Fatalf("HTML should not be escaped: %s", b)
}
if bytes.Contains(b, []byte("你好")) == false {
t.Fatalf("non-ASCII should remain: %s", b)
}
if len(b) == 0 || b[len(b)-1] == '\n' {
t.Fatalf("must not end with newline: %q", b)
}
n, err := protocol.FrameBytes(msg)
if err != nil {
t.Fatal(err)
}
if n != len(b) {
t.Fatalf("FrameBytes=%d len=%d", n, len(b))
}
}
func TestDecodeRoundTripAllFrames(t *testing.T) {
lim := protocol.DefaultLimits()
cases := []struct {
name string
in any
typ string
val func(any) error
}{
{
name: "hello",
in: &protocol.Hello{
V: protocol.Version, Type: protocol.TypeHello, RID: "1",
MaxReceiveBytes: ptrInt(4096), Client: "go-sdk/0.1",
},
typ: protocol.TypeHello,
val: func(v any) error { return v.(*protocol.Hello).Validate() },
},
{
name: "resp_ok",
in: &protocol.Resp{
V: protocol.Version, Type: protocol.TypeResp, RID: "1", OK: true,
Data: protocol.MustRaw(protocol.HelloData{
ServerTimeMs: 1, ServerVersion: "0.1.0",
MaxBodyBytes: 262144, MaxMetaBytes: 4096, MaxFrameBytes: 786432,
MaxTTLSeconds: 1, MaxScheduleSeconds: 1, AckTimeoutSeconds: 1,
SessionToken: "nst_abc",
}),
},
typ: protocol.TypeResp,
val: func(v any) error { return v.(*protocol.Resp).Validate() },
},
{
name: "resp_err",
in: &protocol.Resp{
V: protocol.Version, Type: protocol.TypeResp, RID: "1", OK: false,
Error: &protocol.ErrorBody{Code: protocol.CodeTalkPasswordRequired, Message: "需要对话密码"},
},
typ: protocol.TypeResp,
val: func(v any) error { return v.(*protocol.Resp).Validate() },
},
{
name: "send",
in: &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "2", ID: "018f-Ab",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"},
Meta: map[string]any{"k": "v"},
},
typ: protocol.TypeSend,
val: func(v any) error { return v.(*protocol.Send).Validate(lim) },
},
{
name: "msg",
in: &protocol.Msg{
V: protocol.Version, Type: protocol.TypeMsg, ID: "018f",
From: "app-1", To: protocol.Target{Kind: protocol.TargetGroup, ID: "g_ab12cd34"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"}, SendAtMs: 1,
},
typ: protocol.TypeMsg,
val: func(v any) error { return v.(*protocol.Msg).Validate(lim) },
},
{
name: "ack",
in: &protocol.Ack{V: protocol.Version, Type: protocol.TypeAck, RID: "3", From: "app-1", ID: "018f"},
typ: protocol.TypeAck,
val: func(v any) error { return v.(*protocol.Ack).Validate() },
},
{
name: "recall",
in: &protocol.Recall{V: protocol.Version, Type: protocol.TypeRecall, RID: "4", ID: "018f"},
typ: protocol.TypeRecall,
val: func(v any) error { return v.(*protocol.Recall).Validate() },
},
{
name: "status",
in: &protocol.Status{V: protocol.Version, Type: protocol.TypeStatus, RID: "5", ID: "018f", Limit: 100},
typ: protocol.TypeStatus,
val: func(v any) error { return v.(*protocol.Status).Validate() },
},
{
name: "receipt",
in: &protocol.Receipt{
V: protocol.Version, Type: protocol.TypeReceipt, ReceiptID: "9001",
ID: "018f", EndpointID: "device-1", State: "accepted", AtMs: 1,
},
typ: protocol.TypeReceipt,
val: func(v any) error { return v.(*protocol.Receipt).Validate() },
},
{
name: "receipt_ack",
in: &protocol.ReceiptAck{V: protocol.Version, Type: protocol.TypeReceiptAck, RID: "6", ReceiptID: "9001"},
typ: protocol.TypeReceiptAck,
val: func(v any) error { return v.(*protocol.ReceiptAck).Validate() },
},
{
name: "revoked",
in: &protocol.Revoked{V: protocol.Version, Type: protocol.TypeRevoked, ID: "018f", From: "app-1", Reason: "recalled"},
typ: protocol.TypeRevoked,
val: func(v any) error { return v.(*protocol.Revoked).Validate() },
},
{
name: "presence.get",
in: &protocol.PresenceGet{V: protocol.Version, Type: protocol.TypePresenceGet, RID: "7", IDs: []string{"a", "b"}},
typ: protocol.TypePresenceGet,
val: func(v any) error { return v.(*protocol.PresenceGet).Validate() },
},
{
name: "directory.list",
in: &protocol.DirectoryList{V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "8", Limit: 100},
typ: protocol.TypeDirectoryList,
val: func(v any) error { return v.(*protocol.DirectoryList).Validate() },
},
{
name: "presence.watch",
in: &protocol.PresenceWatch{V: protocol.Version, Type: protocol.TypePresenceWatch, RID: "9", IDs: []string{"a"}, All: false},
typ: protocol.TypePresenceWatch,
val: func(v any) error { return v.(*protocol.PresenceWatch).Validate() },
},
{
name: "presence",
in: &protocol.Presence{V: protocol.Version, Type: protocol.TypePresence, ID: "a", Online: true, AtMs: 1},
typ: protocol.TypePresence,
val: func(v any) error { return v.(*protocol.Presence).Validate() },
},
{
name: "unlock",
in: &protocol.Unlock{V: protocol.Version, Type: protocol.TypeUnlock, RID: "10", EndpointID: "b", TalkPassword: "secret"},
typ: protocol.TypeUnlock,
val: func(v any) error { return v.(*protocol.Unlock).Validate() },
},
{
name: "self.get",
in: &protocol.SelfGet{V: protocol.Version, Type: protocol.TypeSelfGet, RID: "11"},
typ: protocol.TypeSelfGet,
val: func(v any) error { return v.(*protocol.SelfGet).Validate() },
},
{
name: "self.update",
in: &protocol.SelfUpdate{V: protocol.Version, Type: protocol.TypeSelfUpdate, RID: "12", Name: "门口", DefaultDelayMs: ptrInt64(10000)},
typ: protocol.TypeSelfUpdate,
val: func(v any) error { return v.(*protocol.SelfUpdate).Validate() },
},
{
name: "self.talk_password",
in: &protocol.SelfTalkPassword{V: protocol.Version, Type: protocol.TypeSelfTalkPassword, RID: "13", TalkPassword: ""},
typ: protocol.TypeSelfTalkPassword,
val: func(v any) error { return v.(*protocol.SelfTalkPassword).Validate() },
},
{
name: "self.login_password",
in: &protocol.SelfLoginPassword{
V: protocol.Version, Type: protocol.TypeSelfLoginPassword, RID: "14",
OldPassword: "oldpass12", NewPassword: "newpass12",
},
typ: protocol.TypeSelfLoginPassword,
val: func(v any) error { return v.(*protocol.SelfLoginPassword).Validate() },
},
{
name: "self.logout",
in: &protocol.SelfLogout{V: protocol.Version, Type: protocol.TypeSelfLogout, RID: "24"},
typ: protocol.TypeSelfLogout,
val: func(v any) error { return v.(*protocol.SelfLogout).Validate() },
},
{
name: "group.create",
in: &protocol.GroupCreate{
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "15", ID: "",
Name: "一组", Members: []protocol.GroupMemberIn{{ID: "b", TalkPassword: "secret"}},
},
typ: protocol.TypeGroupCreate,
val: func(v any) error { return v.(*protocol.GroupCreate).Validate() },
},
{
name: "group.add",
in: &protocol.GroupAdd{
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "16", GroupID: "g_ab12cd34",
Members: []protocol.GroupMemberIn{{ID: "c"}},
},
typ: protocol.TypeGroupAdd,
val: func(v any) error { return v.(*protocol.GroupAdd).Validate() },
},
{
name: "group.remove",
in: &protocol.GroupRemove{V: protocol.Version, Type: protocol.TypeGroupRemove, RID: "17", GroupID: "g_ab12cd34", EndpointID: "c"},
typ: protocol.TypeGroupRemove,
val: func(v any) error { return v.(*protocol.GroupRemove).Validate() },
},
{
name: "group.leave",
in: &protocol.GroupLeave{V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "18", GroupID: "g_ab12cd34"},
typ: protocol.TypeGroupLeave,
val: func(v any) error { return v.(*protocol.GroupLeave).Validate() },
},
{
name: "group.transfer",
in: &protocol.GroupTransfer{V: protocol.Version, Type: protocol.TypeGroupTransfer, RID: "19", GroupID: "g_ab12cd34", EndpointID: "b"},
typ: protocol.TypeGroupTransfer,
val: func(v any) error { return v.(*protocol.GroupTransfer).Validate() },
},
{
name: "group.rename",
in: &protocol.GroupRename{V: protocol.Version, Type: protocol.TypeGroupRename, RID: "20", GroupID: "g_ab12cd34", Name: "新名"},
typ: protocol.TypeGroupRename,
val: func(v any) error { return v.(*protocol.GroupRename).Validate() },
},
{
name: "group.dissolve",
in: &protocol.GroupDissolve{V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "21", GroupID: "g_ab12cd34"},
typ: protocol.TypeGroupDissolve,
val: func(v any) error { return v.(*protocol.GroupDissolve).Validate() },
},
{
name: "group.list",
in: &protocol.GroupList{V: protocol.Version, Type: protocol.TypeGroupList, RID: "22", Limit: 100},
typ: protocol.TypeGroupList,
val: func(v any) error { return v.(*protocol.GroupList).Validate() },
},
{
name: "group.get",
in: &protocol.GroupGet{V: protocol.Version, Type: protocol.TypeGroupGet, RID: "23", GroupID: "g_ab12cd34", Limit: 100},
typ: protocol.TypeGroupGet,
val: func(v any) error { return v.(*protocol.GroupGet).Validate() },
},
{
name: "group_event",
in: &protocol.GroupEvent{
V: protocol.Version, Type: protocol.TypeGroupEvent, GroupID: "g_ab12cd34",
Event: "member_added", EndpointID: "c", AtMs: 1,
},
typ: protocol.TypeGroupEvent,
val: func(v any) error { return v.(*protocol.GroupEvent).Validate() },
},
{
name: "fatal",
in: &protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: "disabled"},
typ: protocol.TypeFatal,
val: func(v any) error { return v.(*protocol.Fatal).Validate() },
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
raw := mustMarshal(t, tc.in)
got, err := protocol.Decode(raw)
if err != nil {
t.Fatalf("Decode: %v", err)
}
back := mustMarshal(t, got)
var a, b any
if err := json.Unmarshal(raw, &a); err != nil {
t.Fatal(err)
}
if err := json.Unmarshal(back, &b); err != nil {
t.Fatal(err)
}
aj, _ := json.Marshal(a)
bj, _ := json.Marshal(b)
if !bytes.Equal(aj, bj) {
t.Fatalf("round-trip mismatch\n%s\n%s", aj, bj)
}
peek := struct {
Type string `json:"type"`
}{}
_ = json.Unmarshal(raw, &peek)
if peek.Type != tc.typ {
t.Fatalf("type=%s want %s", peek.Type, tc.typ)
}
if err := tc.val(got); err != nil {
t.Fatalf("Validate: %v", err)
}
})
}
}
func TestRegisterRoundTrip(t *testing.T) {
req := &protocol.RegisterRequest{
RegistrationCode: "code",
ID: "device-1",
LoginPassword: "password1",
Name: "门口",
TalkPassword: "talk",
}
raw := mustMarshal(t, req)
got, err := protocol.DecodeRegister(raw)
if err != nil {
t.Fatal(err)
}
if err := got.Validate(); err != nil {
t.Fatal(err)
}
resp := protocol.RegisterResponse{OK: true, Data: protocol.RegisterData{ID: "e_ab12cd34", LoginPassword: "generated"}}
rb := mustMarshal(t, resp)
var decoded protocol.RegisterResponse
if err := protocol.Unmarshal(rb, &decoded); err != nil {
t.Fatal(err)
}
if !decoded.OK || decoded.Data.ID != "e_ab12cd34" {
t.Fatalf("unexpected resp: %+v", decoded)
}
}
func TestValidateTable(t *testing.T) {
lim := protocol.DefaultLimits()
cases := []struct {
name string
check func() error
wantCode string
}{
{
name: "endpoint_id_uppercase",
check: func() error {
return (&protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "Device"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
}).Validate(lim)
},
wantCode: protocol.CodeBadRequest,
},
{
name: "message_id_allows_upper",
check: func() error {
return (&protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "Msg_1.A-b",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
}).Validate(lim)
},
},
{
name: "send_at_and_delay_mutex",
check: func() error {
return (&protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
SendAtMs: ptrInt64(1), DelayMs: ptrInt64(2),
}).Validate(lim)
},
wantCode: protocol.CodeBadRequest,
},
{
name: "body_too_large",
check: func() error {
return (&protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: strings.Repeat("a", 10)},
}).Validate(protocol.Limits{MaxBodyBytes: 5, MaxMetaBytes: 4096, MaxFrameBytes: 786432})
},
wantCode: protocol.CodeBodyTooLarge,
},
{
name: "meta_too_large",
check: func() error {
return (&protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
Meta: map[string]any{"k": strings.Repeat("v", 100)},
}).Validate(protocol.Limits{MaxBodyBytes: 262144, MaxMetaBytes: 20, MaxFrameBytes: 786432})
},
wantCode: protocol.CodeMetaTooLarge,
},
{
name: "login_password_nst_prefix",
check: func() error {
return (&protocol.RegisterRequest{LoginPassword: "nst_notallowed"}).Validate()
},
wantCode: protocol.CodeBadRequest,
},
{
name: "self_login_password_nst_prefix",
check: func() error {
return (&protocol.SelfLoginPassword{
V: protocol.Version, Type: protocol.TypeSelfLoginPassword, RID: "1",
OldPassword: "oldpass12", NewPassword: "nst_tokenlike",
}).Validate()
},
wantCode: protocol.CodeBadRequest,
},
{
name: "hello_max_receive_too_small",
check: func() error {
return (&protocol.Hello{
V: protocol.Version, Type: protocol.TypeHello, RID: "1",
MaxReceiveBytes: ptrInt(100),
}).Validate()
},
wantCode: protocol.CodeBadRequest,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
err := tc.check()
if tc.wantCode == "" {
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
return
}
pe, ok := err.(*protocol.Error)
if !ok {
t.Fatalf("want *protocol.Error, got %T %v", err, err)
}
if pe.Code != tc.wantCode {
t.Fatalf("code=%s want %s", pe.Code, tc.wantCode)
}
})
}
}
func TestRequestFingerprintMetaKeyOrderInsensitive(t *testing.T) {
base := func(meta map[string]any) *protocol.Send {
return &protocol.Send{
V: protocol.Version, Type: protocol.TypeSend, RID: "rid-ignored", ID: "m1",
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
Body: protocol.Body{Enc: protocol.EncUTF8, ContentType: "text/plain", Data: "hello"},
Meta: meta,
Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: ptrInt64(60)},
Receipt: ptrBool(true),
TalkPassword: "should-not-affect",
}
}
fp1, err := protocol.RequestFingerprint(base(map[string]any{"a": 1, "b": "x", "c": true}))
if err != nil {
t.Fatal(err)
}
fp2, err := protocol.RequestFingerprint(base(map[string]any{"c": true, "b": "x", "a": 1}))
if err != nil {
t.Fatal(err)
}
if fp1 != fp2 {
t.Fatalf("fingerprint depends on meta key order: %s vs %s", fp1, fp2)
}
// rid / talk_password 变化不应影响
s3 := base(map[string]any{"a": 1, "b": "x", "c": true})
s3.RID = "other-rid"
s3.TalkPassword = "other"
fp3, err := protocol.RequestFingerprint(s3)
if err != nil {
t.Fatal(err)
}
if fp3 != fp1 {
t.Fatalf("rid/talk_password affected fingerprint")
}
// 正文变化应影响
s4 := base(map[string]any{"a": 1, "b": "x", "c": true})
s4.Body.Data = "hello!"
fp4, err := protocol.RequestFingerprint(s4)
if err != nil {
t.Fatal(err)
}
if fp4 == fp1 {
t.Fatalf("body change should change fingerprint")
}
}
func TestErrorCodesDefined(t *testing.T) {
codes := []string{
protocol.CodeBadRequest, protocol.CodeNotReady, protocol.CodeUnauthorized,
protocol.CodeForbidden, protocol.CodeNotFound, protocol.CodeInvalidTarget,
protocol.CodeConflict, protocol.CodeIDTaken, protocol.CodeBodyTooLarge,
protocol.CodeMetaTooLarge, protocol.CodeFrameTooLarge, protocol.CodeResponseTooLarge,
protocol.CodeTalkPasswordRequired, protocol.CodeTalkPasswordInvalid,
protocol.CodeRateLimited, protocol.CodeNotMember, protocol.CodeOwnerCannotLeave,
protocol.CodeGroupFull, protocol.CodeQuotaExceeded, protocol.CodeEndpointDisabled,
protocol.CodeRegistrationClosed, protocol.CodeRegistrationCodeInvalid, protocol.CodeBusy,
}
if len(codes) != 23 {
t.Fatalf("want 23 error codes, got %d", len(codes))
}
seen := map[string]bool{}
for _, c := range codes {
if c == "" || seen[c] {
t.Fatalf("bad code %q", c)
}
seen[c] = true
}
}