diff --git a/internal/protocol/const.go b/internal/protocol/const.go new file mode 100644 index 0000000..b4c396c --- /dev/null +++ b/internal/protocol/const.go @@ -0,0 +1,94 @@ +package protocol + +// 协议常量与默认上限(与 DEVELOPMENT 第 6、11 节示例一致)。 + +const ( + Version = 1 + + TypeHello = "hello" + TypeResp = "resp" + TypeSend = "send" + TypeMsg = "msg" + TypeAck = "ack" + TypeRecall = "recall" + TypeStatus = "status" + TypeReceipt = "receipt" + TypeReceiptAck = "receipt_ack" + TypeRevoked = "revoked" + TypePresenceGet = "presence.get" + TypeDirectoryList = "directory.list" + TypePresenceWatch = "presence.watch" + TypePresence = "presence" + TypeUnlock = "unlock" + TypeSelfGet = "self.get" + TypeSelfUpdate = "self.update" + TypeSelfTalkPassword = "self.talk_password" + TypeSelfLoginPassword = "self.login_password" + TypeSelfLogout = "self.logout" + TypeGroupCreate = "group.create" + TypeGroupAdd = "group.add" + TypeGroupRemove = "group.remove" + TypeGroupLeave = "group.leave" + TypeGroupTransfer = "group.transfer" + TypeGroupRename = "group.rename" + TypeGroupDissolve = "group.dissolve" + TypeGroupList = "group.list" + TypeGroupGet = "group.get" + TypeGroupEvent = "group_event" + TypeFatal = "fatal" + + TargetEndpoint = "endpoint" + TargetGroup = "group" + + EncUTF8 = "utf8" + EncBase64 = "base64" + + DefaultContentTypeUTF8 = "text/plain; charset=utf-8" + DefaultContentTypeBase64 = "application/octet-stream" + + SessionTokenPrefix = "nst_" + + MinIDLen = 1 + MaxIDLen = 64 + MaxContentTypeLen = 128 + MaxNameChars = 64 + MinLoginPasswordLen = 8 + MaxLoginPasswordLen = 128 + MinTalkPasswordLen = 4 + MaxTalkPasswordLen = 64 + MinMaxReceiveBytes = 1024 + MaxPresenceGetIDs = 200 + MaxPresenceWatchIDs = 1000 + MaxPageLimit = 200 + DefaultOfflineTTL = int64(86400) + DefaultMaxBodyBytes = 262144 + DefaultMaxMetaBytes = 4096 + DefaultMaxFrameBytes = 786432 +) + +// 第 6.10 节错误码。 +const ( + CodeBadRequest = "bad_request" + CodeNotReady = "not_ready" + CodeUnauthorized = "unauthorized" + CodeForbidden = "forbidden" + CodeNotFound = "not_found" + CodeInvalidTarget = "invalid_target" + CodeConflict = "conflict" + CodeIDTaken = "id_taken" + CodeBodyTooLarge = "body_too_large" + CodeMetaTooLarge = "meta_too_large" + CodeFrameTooLarge = "frame_too_large" + CodeResponseTooLarge = "response_too_large" + CodeTalkPasswordRequired = "talk_password_required" + CodeTalkPasswordInvalid = "talk_password_invalid" + CodeRateLimited = "rate_limited" + CodeNotMember = "not_member" + CodeOwnerCannotLeave = "owner_cannot_leave" + CodeGroupFull = "group_full" + CodeQuotaExceeded = "quota_exceeded" + CodeEndpointDisabled = "endpoint_disabled" + CodeRegistrationClosed = "registration_closed" + CodeRegistrationCodeInvalid = "registration_code_invalid" + CodeBusy = "busy" +) diff --git a/internal/protocol/decode.go b/internal/protocol/decode.go new file mode 100644 index 0000000..7116497 --- /dev/null +++ b/internal/protocol/decode.go @@ -0,0 +1,115 @@ +package protocol + +import ( + "encoding/json" + "fmt" +) + +type typePeek struct { + V int `json:"v"` + Type string `json:"type"` +} + +// Decode 解析一帧 MQTT/应用 JSON,返回具体类型。 +func Decode(data []byte) (any, error) { + var peek typePeek + if err := Unmarshal(data, &peek); err != nil { + return nil, badRequest("invalid json") + } + if peek.V != 0 && peek.V != Version { + return nil, badRequest("unsupported version") + } + if peek.Type == "" { + return nil, badRequest("missing type") + } + + var out any + switch peek.Type { + case TypeHello: + out = &Hello{} + case TypeResp: + out = &Resp{} + case TypeSend: + out = &Send{} + case TypeMsg: + out = &Msg{} + case TypeAck: + out = &Ack{} + case TypeRecall: + out = &Recall{} + case TypeStatus: + out = &Status{} + case TypeReceipt: + out = &Receipt{} + case TypeReceiptAck: + out = &ReceiptAck{} + case TypeRevoked: + out = &Revoked{} + case TypePresenceGet: + out = &PresenceGet{} + case TypeDirectoryList: + out = &DirectoryList{} + case TypePresenceWatch: + out = &PresenceWatch{} + case TypePresence: + out = &Presence{} + case TypeUnlock: + out = &Unlock{} + case TypeSelfGet: + out = &SelfGet{} + case TypeSelfUpdate: + out = &SelfUpdate{} + case TypeSelfTalkPassword: + out = &SelfTalkPassword{} + case TypeSelfLoginPassword: + out = &SelfLoginPassword{} + case TypeSelfLogout: + out = &SelfLogout{} + case TypeGroupCreate: + out = &GroupCreate{} + case TypeGroupAdd: + out = &GroupAdd{} + case TypeGroupRemove: + out = &GroupRemove{} + case TypeGroupLeave: + out = &GroupLeave{} + case TypeGroupTransfer: + out = &GroupTransfer{} + case TypeGroupRename: + out = &GroupRename{} + case TypeGroupDissolve: + out = &GroupDissolve{} + case TypeGroupList: + out = &GroupList{} + case TypeGroupGet: + out = &GroupGet{} + case TypeGroupEvent: + out = &GroupEvent{} + case TypeFatal: + out = &Fatal{} + default: + return nil, badRequest(fmt.Sprintf("unknown type %q", peek.Type)) + } + if err := Unmarshal(data, out); err != nil { + return nil, badRequest("invalid json for type") + } + return out, nil +} + +// DecodeRegister 解析注册 HTTP 请求体。 +func DecodeRegister(data []byte) (*RegisterRequest, error) { + var req RegisterRequest + if err := Unmarshal(data, &req); err != nil { + return nil, badRequest("invalid json") + } + return &req, nil +} + +// MustRaw 将值编码为 json.RawMessage(用于填 Resp.Data)。 +func MustRaw(v any) json.RawMessage { + b, err := Marshal(v) + if err != nil { + panic(err) + } + return b +} diff --git a/internal/protocol/encode.go b/internal/protocol/encode.go new file mode 100644 index 0000000..95b6e3f --- /dev/null +++ b/internal/protocol/encode.go @@ -0,0 +1,48 @@ +package protocol + +import ( + "bytes" + "encoding/json" + "io" +) + +// Marshal 将值编码为 UTF-8 JSON,不转义 HTML 与非 ASCII,且不含尾部换行。 +func Marshal(v any) ([]byte, error) { + var buf bytes.Buffer + if err := Encode(&buf, v); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +// Encode 写入 JSON。json.Encoder 默认会追加 '\n',这里去掉,保证整帧字节数与线上一致。 +func Encode(w io.Writer, v any) error { + var buf bytes.Buffer + enc := json.NewEncoder(&buf) + enc.SetEscapeHTML(false) + if err := enc.Encode(v); err != nil { + return err + } + b := buf.Bytes() + if n := len(b); n > 0 && b[n-1] == '\n' { + b = b[:n-1] + } + _, err := w.Write(b) + return err +} + +// Unmarshal 解码 JSON。 +func Unmarshal(data []byte, v any) error { + dec := json.NewDecoder(bytes.NewReader(data)) + dec.UseNumber() + return dec.Decode(v) +} + +// FrameBytes 返回整帧编码后的字节数。 +func FrameBytes(v any) (int, error) { + b, err := Marshal(v) + if err != nil { + return 0, err + } + return len(b), nil +} diff --git a/internal/protocol/error.go b/internal/protocol/error.go new file mode 100644 index 0000000..9c7c819 --- /dev/null +++ b/internal/protocol/error.go @@ -0,0 +1,46 @@ +package protocol + +// ErrorBody 是 resp / 注册失败里的 error 对象。 +type ErrorBody struct { + Code string `json:"code"` + Message string `json:"message"` +} + +// Error 是协议包校验失败时返回的错误,Code 对应第 6.10 节。 +type Error struct { + Code string + Message string +} + +func (e *Error) Error() string { + if e == nil { + return "" + } + if e.Message == "" { + return e.Code + } + return e.Code + ": " + e.Message +} + +func (e *Error) Body() ErrorBody { + if e == nil { + return ErrorBody{} + } + return ErrorBody{Code: e.Code, Message: e.Message} +} + +func badRequest(msg string) *Error { + return &Error{Code: CodeBadRequest, Message: msg} +} + +func bodyTooLarge(msg string) *Error { + return &Error{Code: CodeBodyTooLarge, Message: msg} +} + +func metaTooLarge(msg string) *Error { + return &Error{Code: CodeMetaTooLarge, Message: msg} +} + +func frameTooLarge(msg string) *Error { + return &Error{Code: CodeFrameTooLarge, Message: msg} +} diff --git a/internal/protocol/fingerprint.go b/internal/protocol/fingerprint.go new file mode 100644 index 0000000..f5eb284 --- /dev/null +++ b/internal/protocol/fingerprint.go @@ -0,0 +1,81 @@ +package protocol + +import ( + "crypto/sha256" + "encoding/binary" + "encoding/hex" +) + +// RequestFingerprint 计算第 7.3 节发送请求指纹(十六进制小写 SHA-256)。 +// 字段:to.kind、to.id、body.enc、content_type、解码后正文、meta(键排序)、 +// send_at_ms、delay_ms、offline.keep、offline.ttl_seconds、receipt。 +// 不含 talk_password 与 rid。meta 键顺序不影响结果。 +// receipt / offline 使用文档默认值后的有效值;缺省的 send_at_ms、delay_ms 不写入对应槽位。 +func RequestFingerprint(s *Send) (string, error) { + raw, err := DecodeBody(s.Body) + if err != nil { + return "", err + } + metaJSON, err := MetaCanonicalJSON(s.Meta) + if err != nil { + return "", badRequest("invalid meta") + } + + h := sha256.New() + writeStr(h, s.To.Kind) + writeStr(h, s.To.ID) + writeStr(h, s.Body.Enc) + writeStr(h, EffectiveContentType(s.Body)) + writeBytes(h, raw) + writeBytes(h, metaJSON) + + writeOptInt64(h, s.SendAtMs) + writeOptInt64(h, s.DelayMs) + + keep := EffectiveOfflineKeep(s) + writeBool(h, keep) + ttl := EffectiveOfflineTTL(s) + writeInt64(h, ttl) + writeBool(h, EffectiveReceipt(s)) + + sum := h.Sum(nil) + return hex.EncodeToString(sum), nil +} + +type hashWriter interface { + Write(p []byte) (int, error) +} + +func writeStr(h hashWriter, s string) { + writeBytes(h, []byte(s)) +} + +func writeBytes(h hashWriter, b []byte) { + var lenBuf [8]byte + binary.BigEndian.PutUint64(lenBuf[:], uint64(len(b))) + _, _ = h.Write(lenBuf[:]) + _, _ = h.Write(b) +} + +func writeBool(h hashWriter, v bool) { + if v { + _, _ = h.Write([]byte{1}) + } else { + _, _ = h.Write([]byte{0}) + } +} + +func writeInt64(h hashWriter, v int64) { + var buf [8]byte + binary.BigEndian.PutUint64(buf[:], uint64(v)) + _, _ = h.Write(buf[:]) +} + +func writeOptInt64(h hashWriter, v *int64) { + if v == nil { + _, _ = h.Write([]byte{0}) + return + } + _, _ = h.Write([]byte{1}) + writeInt64(h, *v) +} diff --git a/internal/protocol/ids.go b/internal/protocol/ids.go new file mode 100644 index 0000000..37e3337 --- /dev/null +++ b/internal/protocol/ids.go @@ -0,0 +1,94 @@ +package protocol + +import ( + "strings" + "unicode/utf8" +) + +func isEndpointIDChar(c rune) bool { + return (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_' || c == '.' || c == '-' +} + +func isMessageIDChar(c rune) bool { + return (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '_' || c == '.' || c == '-' +} + +// ValidEndpointID 校验端编号或群编号(小写字母、数字、_ . -,1–64)。 +func ValidEndpointID(id string) bool { + n := len(id) + if n < MinIDLen || n > MaxIDLen { + return false + } + for _, c := range id { + if !isEndpointIDChar(c) { + return false + } + } + return true +} + +// ValidMessageID 校验消息号(字母数字 _ . -,区分大小写,1–64)。 +func ValidMessageID(id string) bool { + n := len(id) + if n < MinIDLen || n > MaxIDLen { + return false + } + for _, c := range id { + if !isMessageIDChar(c) { + return false + } + } + return true +} + +// ValidLoginPassword 校验登录密码:8–128 字符,且不能以 nst_ 开头。空串表示由服务器生成,视为合法。 +func ValidLoginPassword(pw string) bool { + if pw == "" { + return true + } + if strings.HasPrefix(pw, SessionTokenPrefix) { + return false + } + n := utf8.RuneCountInString(pw) + return n >= MinLoginPasswordLen && n <= MaxLoginPasswordLen +} + +// LoginPasswordForbiddenPrefix 报告密码是否因 nst_ 前缀非法。 +func LoginPasswordForbiddenPrefix(pw string) bool { + return strings.HasPrefix(pw, SessionTokenPrefix) +} + +// ValidTalkPassword 校验对话密码:空表示清除/不设;否则 4–64 字符。 +func ValidTalkPassword(pw string) bool { + if pw == "" { + return true + } + n := utf8.RuneCountInString(pw) + return n >= MinTalkPasswordLen && n <= MaxTalkPasswordLen +} + +// ValidName 校验名称:最多 64 个 Unicode 字符。 +func ValidName(name string) bool { + return utf8.RuneCountInString(name) <= MaxNameChars +} + +func requireRID(rid string) error { + if rid == "" { + return badRequest("missing rid") + } + return nil +} + +func requireVersion(v int) error { + if v != Version { + return badRequest("unsupported version") + } + return nil +} + +func requireType(got, want string) error { + if got != want { + return badRequest("wrong type") + } + return nil +} diff --git a/internal/protocol/protocol_test.go b/internal/protocol/protocol_test.go new file mode 100644 index 0000000..82f0f23 --- /dev/null +++ b/internal/protocol/protocol_test.go @@ -0,0 +1,535 @@ +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: "你好