package message import ( "context" "encoding/json" "errors" "sync" "git.asio.asia/nixevol/NixMsg/internal/app/port" ) var ( errPayloadTooLarge = errors.New("message: payload too large") errPublishFailed = errors.New("message: publish failed") ) // LiveConn 是端当前连接的快照(含握手中)。 type LiveConn struct { ConnID port.ConnID MaxReceiveBytes int MaxPacketSize uint32 } // ConnRegistry 查询端是否有连接(由 N 线或测试假实现注入)。 // 有连接(含握手中)即视为在线,用于分发时计算 expire_at。 type ConnRegistry interface { Current(endpointID string) (LiveConn, bool) } // MemoryConns 是测试用的内存连接表。 type MemoryConns struct { mu sync.RWMutex m map[string]LiveConn } // NewMemoryConns 创建空连接表。 func NewMemoryConns() *MemoryConns { return &MemoryConns{m: make(map[string]LiveConn)} } // Set 登记或更新端的当前连接。 func (c *MemoryConns) Set(endpointID string, conn LiveConn) { c.mu.Lock() defer c.mu.Unlock() c.m[endpointID] = conn } // Clear 移除端的当前连接;若代号不匹配则不动。 func (c *MemoryConns) Clear(endpointID string, connID port.ConnID) { c.mu.Lock() defer c.mu.Unlock() cur, ok := c.m[endpointID] if !ok { return } if connID != "" && cur.ConnID != connID { return } delete(c.m, endpointID) } // Current 实现 ConnRegistry。 func (c *MemoryConns) Current(endpointID string) (LiveConn, bool) { c.mu.RLock() defer c.mu.RUnlock() v, ok := c.m[endpointID] return v, ok } // RecordingDownlink 记录下行发布,供测试断言。 type RecordingDownlink struct { mu sync.Mutex Published []DownPublish FailNext int // 接下来 N 次 PublishDown 返回错误 MaxSize int // >0 时超限返回错误 } // DownPublish 是一次下行记录。 type DownPublish struct { EndpointID string ConnID port.ConnID Payload []byte QoS byte } // PublishDown 实现 port.Downlink。 func (d *RecordingDownlink) PublishDown(_ context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { d.mu.Lock() defer d.mu.Unlock() if d.MaxSize > 0 && len(payload) > d.MaxSize { return errPayloadTooLarge } if d.FailNext > 0 { d.FailNext-- return errPublishFailed } d.Published = append(d.Published, DownPublish{ EndpointID: endpointID, ConnID: connID, Payload: append([]byte(nil), payload...), QoS: opts.QoS, }) return nil } // Count 返回已发布条数。 func (d *RecordingDownlink) Count() int { d.mu.Lock() defer d.mu.Unlock() return len(d.Published) } // Snapshots 返回发布副本。 func (d *RecordingDownlink) Snapshots() []DownPublish { d.mu.Lock() defer d.mu.Unlock() out := make([]DownPublish, len(d.Published)) copy(out, d.Published) return out } // FilterType 统计 type 字段匹配的发布次数。 func (d *RecordingDownlink) FilterType(typ string) int { d.mu.Lock() defer d.mu.Unlock() n := 0 for _, p := range d.Published { if payloadType(p.Payload) == typ { n++ } } return n } func payloadType(payload []byte) string { var head struct { Type string `json:"type"` } _ = json.Unmarshal(payload, &head) return head.Type } var _ port.Downlink = (*RecordingDownlink)(nil) var _ ConnRegistry = (*MemoryConns)(nil)