package protocol import ( "encoding/base64" "encoding/json" "unicode/utf8" ) // DecodeBody 按 enc 解码正文,返回原始字节。 func DecodeBody(b Body) ([]byte, error) { switch b.Enc { case EncUTF8: if !utf8.ValidString(b.Data) { return nil, badRequest("body data is not valid utf8") } return []byte(b.Data), nil case EncBase64: raw, err := base64.StdEncoding.DecodeString(b.Data) if err != nil { return nil, badRequest("body data is not valid base64") } return raw, nil default: return nil, badRequest("invalid body.enc") } } // EffectiveContentType 返回带默认值的 content_type。 func EffectiveContentType(b Body) string { if b.ContentType != "" { return b.ContentType } switch b.Enc { case EncBase64: return DefaultContentTypeBase64 default: return DefaultContentTypeUTF8 } } // EffectiveReceipt 返回 receipt 有效值(默认 true)。 func EffectiveReceipt(s *Send) bool { if s.Receipt == nil { return true } return *s.Receipt } // EffectiveOfflineKeep 返回 offline.keep(默认 false)。 func EffectiveOfflineKeep(s *Send) bool { if s.Offline == nil { return false } return s.Offline.Keep } // EffectiveOfflineTTL 返回 offline.ttl_seconds 有效值;keep 为 false 时为 0。 func EffectiveOfflineTTL(s *Send) int64 { if !EffectiveOfflineKeep(s) { return 0 } if s.Offline.TTLSeconds == nil { return DefaultOfflineTTL } return *s.Offline.TTLSeconds } func checkMetaSize(meta map[string]any, maxBytes int) error { if meta == nil { return nil } b, err := Marshal(meta) if err != nil { return badRequest("invalid meta") } if len(b) > maxBytes { return metaTooLarge("meta exceeds limit") } return nil } func checkBodySize(b Body, maxBytes int) ([]byte, error) { raw, err := DecodeBody(b) if err != nil { if pe, ok := err.(*Error); ok { return nil, pe } return nil, badRequest(err.Error()) } if len(raw) > maxBytes { return nil, bodyTooLarge("body exceeds limit") } if b.ContentType != "" && utf8.RuneCountInString(b.ContentType) > MaxContentTypeLen { return nil, badRequest("content_type too long") } return raw, nil } func checkPageLimit(limit int) error { if limit < 0 || limit > MaxPageLimit { return badRequest("limit out of range") } return nil } // Validate 校验 Hello。 func (h *Hello) Validate() error { if err := requireVersion(h.V); err != nil { return err } if err := requireType(h.Type, TypeHello); err != nil { return err } if err := requireRID(h.RID); err != nil { return err } if h.MaxReceiveBytes != nil && *h.MaxReceiveBytes < MinMaxReceiveBytes { return badRequest("max_receive_bytes too small") } return nil } // Validate 校验 Resp。 func (r *Resp) Validate() error { if err := requireVersion(r.V); err != nil { return err } if err := requireType(r.Type, TypeResp); err != nil { return err } if err := requireRID(r.RID); err != nil { return err } if !r.OK && (r.Error == nil || r.Error.Code == "") { return badRequest("error response missing code") } return nil } // Validate 校验 Send(含正文/meta 大小与 send_at_ms/delay_ms 互斥)。 func (s *Send) Validate(lim Limits) error { lim = lim.withDefaults() if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeSend); err != nil { return err } if err := requireRID(s.RID); err != nil { return err } if !ValidMessageID(s.ID) { return badRequest("invalid message id") } if s.To.Kind != TargetEndpoint && s.To.Kind != TargetGroup { return badRequest("invalid to.kind") } if !ValidEndpointID(s.To.ID) { return badRequest("invalid to.id") } if _, err := checkBodySize(s.Body, lim.MaxBodyBytes); err != nil { return err } if err := checkMetaSize(s.Meta, lim.MaxMetaBytes); err != nil { return err } if s.SendAtMs != nil && s.DelayMs != nil { return badRequest("send_at_ms and delay_ms are mutually exclusive") } n, err := FrameBytes(s) if err != nil { return badRequest("cannot encode frame") } if n > lim.MaxFrameBytes { return frameTooLarge("frame exceeds limit") } return nil } // Validate 校验 Msg。 func (m *Msg) Validate(lim Limits) error { lim = lim.withDefaults() if err := requireVersion(m.V); err != nil { return err } if err := requireType(m.Type, TypeMsg); err != nil { return err } if !ValidMessageID(m.ID) { return badRequest("invalid message id") } if !ValidEndpointID(m.From) { return badRequest("invalid from") } if m.To.Kind != TargetEndpoint && m.To.Kind != TargetGroup { return badRequest("invalid to.kind") } if !ValidEndpointID(m.To.ID) { return badRequest("invalid to.id") } if _, err := checkBodySize(m.Body, lim.MaxBodyBytes); err != nil { return err } if err := checkMetaSize(m.Meta, lim.MaxMetaBytes); err != nil { return err } return nil } // Validate 校验 Ack。 func (a *Ack) Validate() error { if err := requireVersion(a.V); err != nil { return err } if err := requireType(a.Type, TypeAck); err != nil { return err } if err := requireRID(a.RID); err != nil { return err } if !ValidEndpointID(a.From) { return badRequest("invalid from") } if !ValidMessageID(a.ID) { return badRequest("invalid message id") } return nil } // Validate 校验 Recall。 func (r *Recall) Validate() error { if err := requireVersion(r.V); err != nil { return err } if err := requireType(r.Type, TypeRecall); err != nil { return err } if err := requireRID(r.RID); err != nil { return err } if !ValidMessageID(r.ID) { return badRequest("invalid message id") } return nil } // Validate 校验 Status。 func (s *Status) Validate() error { if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeStatus); err != nil { return err } if err := requireRID(s.RID); err != nil { return err } if !ValidMessageID(s.ID) { return badRequest("invalid message id") } if err := checkPageLimit(s.Limit); err != nil { return err } return nil } // Validate 校验 Receipt。 func (r *Receipt) Validate() error { if err := requireVersion(r.V); err != nil { return err } if err := requireType(r.Type, TypeReceipt); err != nil { return err } if r.ReceiptID == "" { return badRequest("missing receipt_id") } if !ValidMessageID(r.ID) { return badRequest("invalid message id") } if r.EndpointID != "" && !ValidEndpointID(r.EndpointID) { return badRequest("invalid endpoint_id") } return nil } // Validate 校验 ReceiptAck。 func (r *ReceiptAck) Validate() error { if err := requireVersion(r.V); err != nil { return err } if err := requireType(r.Type, TypeReceiptAck); err != nil { return err } if err := requireRID(r.RID); err != nil { return err } if r.ReceiptID == "" { return badRequest("missing receipt_id") } return nil } // Validate 校验 Revoked。 func (r *Revoked) Validate() error { if err := requireVersion(r.V); err != nil { return err } if err := requireType(r.Type, TypeRevoked); err != nil { return err } if !ValidMessageID(r.ID) { return badRequest("invalid message id") } if !ValidEndpointID(r.From) { return badRequest("invalid from") } if r.Reason == "" { return badRequest("missing reason") } return nil } // Validate 校验 PresenceGet。 func (p *PresenceGet) Validate() error { if err := requireVersion(p.V); err != nil { return err } if err := requireType(p.Type, TypePresenceGet); err != nil { return err } if err := requireRID(p.RID); err != nil { return err } if len(p.IDs) > MaxPresenceGetIDs { return badRequest("ids count out of range") } for _, id := range p.IDs { if !ValidEndpointID(id) { return badRequest("invalid id in ids") } } return nil } // Validate 校验 DirectoryList。 func (d *DirectoryList) Validate() error { if err := requireVersion(d.V); err != nil { return err } if err := requireType(d.Type, TypeDirectoryList); err != nil { return err } if err := requireRID(d.RID); err != nil { return err } if err := checkPageLimit(d.Limit); err != nil { return err } return nil } // Validate 校验 PresenceWatch。 func (p *PresenceWatch) Validate() error { if err := requireVersion(p.V); err != nil { return err } if err := requireType(p.Type, TypePresenceWatch); err != nil { return err } if err := requireRID(p.RID); err != nil { return err } if len(p.IDs) > MaxPresenceWatchIDs { return badRequest("ids count out of range") } for _, id := range p.IDs { if !ValidEndpointID(id) { return badRequest("invalid id in ids") } } return nil } // Validate 校验 Presence。 func (p *Presence) Validate() error { if err := requireVersion(p.V); err != nil { return err } if err := requireType(p.Type, TypePresence); err != nil { return err } if !ValidEndpointID(p.ID) { return badRequest("invalid id") } return nil } // Validate 校验 Unlock。 func (u *Unlock) Validate() error { if err := requireVersion(u.V); err != nil { return err } if err := requireType(u.Type, TypeUnlock); err != nil { return err } if err := requireRID(u.RID); err != nil { return err } if !ValidEndpointID(u.EndpointID) { return badRequest("invalid endpoint_id") } return nil } // Validate 校验 SelfGet。 func (s *SelfGet) Validate() error { if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeSelfGet); err != nil { return err } return requireRID(s.RID) } // Validate 校验 SelfUpdate。 func (s *SelfUpdate) Validate() error { if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeSelfUpdate); err != nil { return err } if err := requireRID(s.RID); err != nil { return err } if !ValidName(s.Name) { return badRequest("name too long") } if s.DefaultDelayMs != nil && *s.DefaultDelayMs < 0 { return badRequest("default_delay_ms negative") } return nil } // Validate 校验 SelfTalkPassword。 func (s *SelfTalkPassword) Validate() error { if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeSelfTalkPassword); err != nil { return err } if err := requireRID(s.RID); err != nil { return err } if !ValidTalkPassword(s.TalkPassword) { return badRequest("invalid talk_password") } return nil } // Validate 校验 SelfLoginPassword。 func (s *SelfLoginPassword) Validate() error { if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeSelfLoginPassword); err != nil { return err } if err := requireRID(s.RID); err != nil { return err } if s.NewPassword == "" { return badRequest("new_password required") } if LoginPasswordForbiddenPrefix(s.NewPassword) { return badRequest("login password must not start with nst_") } if !ValidLoginPassword(s.NewPassword) { return badRequest("invalid new_password") } return nil } // Validate 校验 SelfLogout。 func (s *SelfLogout) Validate() error { if err := requireVersion(s.V); err != nil { return err } if err := requireType(s.Type, TypeSelfLogout); err != nil { return err } return requireRID(s.RID) } func validateGroupMembers(members []GroupMemberIn) error { for _, m := range members { if !ValidEndpointID(m.ID) { return badRequest("invalid member id") } } return nil } // Validate 校验 GroupCreate。 func (g *GroupCreate) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupCreate); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if g.ID != "" && !ValidEndpointID(g.ID) { return badRequest("invalid group id") } if !ValidName(g.Name) || g.Name == "" { return badRequest("invalid name") } return validateGroupMembers(g.Members) } // Validate 校验 GroupAdd。 func (g *GroupAdd) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupAdd); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } if len(g.Members) == 0 { return badRequest("members required") } return validateGroupMembers(g.Members) } // Validate 校验 GroupRemove。 func (g *GroupRemove) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupRemove); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } if !ValidEndpointID(g.EndpointID) { return badRequest("invalid endpoint_id") } return nil } // Validate 校验 GroupLeave。 func (g *GroupLeave) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupLeave); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } return nil } // Validate 校验 GroupTransfer。 func (g *GroupTransfer) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupTransfer); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } if !ValidEndpointID(g.EndpointID) { return badRequest("invalid endpoint_id") } return nil } // Validate 校验 GroupRename。 func (g *GroupRename) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupRename); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } if !ValidName(g.Name) || g.Name == "" { return badRequest("invalid name") } return nil } // Validate 校验 GroupDissolve。 func (g *GroupDissolve) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupDissolve); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } return nil } // Validate 校验 GroupList。 func (g *GroupList) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupList); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } return checkPageLimit(g.Limit) } // Validate 校验 GroupGet。 func (g *GroupGet) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupGet); err != nil { return err } if err := requireRID(g.RID); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } return checkPageLimit(g.Limit) } // Validate 校验 GroupEvent。 func (g *GroupEvent) Validate() error { if err := requireVersion(g.V); err != nil { return err } if err := requireType(g.Type, TypeGroupEvent); err != nil { return err } if !ValidEndpointID(g.GroupID) { return badRequest("invalid group_id") } if g.Event == "" { return badRequest("missing event") } if g.EndpointID != "" && !ValidEndpointID(g.EndpointID) { return badRequest("invalid endpoint_id") } return nil } // Validate 校验 Fatal。 func (f *Fatal) Validate() error { if err := requireVersion(f.V); err != nil { return err } if err := requireType(f.Type, TypeFatal); err != nil { return err } switch f.Reason { case "disabled", "deleted", "password_reset", "protocol": return nil default: return badRequest("invalid reason") } } // Validate 校验 RegisterRequest。 func (r *RegisterRequest) Validate() error { if r.ID != "" && !ValidEndpointID(r.ID) { return badRequest("invalid id") } if LoginPasswordForbiddenPrefix(r.LoginPassword) { return badRequest("login password must not start with nst_") } if !ValidLoginPassword(r.LoginPassword) { return badRequest("invalid login_password") } if !ValidName(r.Name) { return badRequest("name too long") } if !ValidTalkPassword(r.TalkPassword) { return badRequest("invalid talk_password") } return nil } // MetaCanonicalJSON 将 meta 规范化后序列化(键排序;嵌套 map 同样处理)。 // encoding/json 对 map[string]T 会按键名字典序输出,因此插入顺序不影响结果。 func MetaCanonicalJSON(meta map[string]any) ([]byte, error) { if meta == nil { return []byte("{}"), nil } return Marshal(normalizeMetaValue(meta)) } func normalizeMetaValue(v any) any { switch t := v.(type) { case map[string]any: out := make(map[string]any, len(t)) for k, child := range t { out[k] = normalizeMetaValue(child) } return out case json.Number: return t default: return v } }