Files
NixMsg/internal/app/message/conn.go
T

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)