Files

751 lines
17 KiB
Go

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
}
}