156 lines
3.5 KiB
Go
156 lines
3.5 KiB
Go
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
|
|
}
|
|
|
|
// Snapshot 返回当前连接表副本(供推送循环遍历)。
|
|
func (c *MemoryConns) Snapshot() map[string]LiveConn {
|
|
c.mu.RLock()
|
|
defer c.mu.RUnlock()
|
|
out := make(map[string]LiveConn, len(c.m))
|
|
for k, v := range c.m {
|
|
out[k] = v
|
|
}
|
|
return out
|
|
}
|
|
|
|
// 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)
|