Files
HomeAgent/internal/plugins/remotedevice/registry.go
JianFeeeee f2e3215c77 feat: remotedevice 二进制分块协议(录像回传+音频下发) + 工具描述增强
回应 GUI c29abe9 需服务端配套项([阻塞]两项 + [中]一项):

1. readFrame 支持 0x2 二进制帧(之前遇 0x2 直接断连):
   - 返回 opcode, 二进制帧上限放宽至 8MB
2. handleWS 二进制聚合协议(设备→网关, 录像回传):
   - cmd_data_start 开启按 req_id 聚合 → 0x2 帧追加 → cmd_data_end 聚合完成
   - 结果存 cmdresult(data_base64 字段), device_ctl_cmdresult 可取回
   - 超限防护(声明 total×2 或硬上限 64MB)
3. PushData 下发协议(网关→设备, 音频/TTS):
   - cmd_speech_start → 0x2 分块(8KB) → cmd_speech_end
   - GUI 端已实现接收侧(speakeruse 播放链路打通)
4. device_ctl_cmdrun 工具描述补齐:
   - screensue 带参示例(<内容>/<秒> <内容>)
   - camerasue 录像说明(cmdresult 含 data_base64)
   - speakeruse 文字朗读

测试: 最小 WS 客户端端到端验证二进制分块上传/PushData 下发/离线报错
2026-08-21 10:53:38 +08:00

733 lines
19 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package remotedevice
import (
"bufio"
"crypto/sha256"
"encoding/base64"
"encoding/binary"
"encoding/json"
"fmt"
"io"
"log"
"net"
"net/http"
"strings"
"sync"
"time"
)
// DeviceMeta 描述一台接入了网关的设备。
type DeviceMeta struct {
DeviceID string `json:"device_id"`
Name string `json:"name"`
Kind string `json:"kind"`
Caps []string `json:"caps"`
Info map[string]interface{} `json:"info,omitempty"`
Authorized bool `json:"authorized"`
Online bool `json:"online"`
LastSeen int64 `json:"last_seen"`
RemoteAddr string `json:"remote_addr"`
}
// wconn 表示一条活跃的 WS 连接(由网关持有)。
type wconn struct {
deviceID string
w *bufio.Writer
}
// Registry 是设备接入网关的注册表:管理在线连接、设备元数据与已授权集合。线程安全。
type Registry struct {
mu sync.RWMutex
devices map[string]*DeviceMeta // deviceID -> meta在线/历史)
authorized map[string]bool // deviceID -> 是否已授权(持久化恢复)
conns map[string]*wconn // deviceID -> 活跃连接(支持 push
onlineCh chan string
onStatus func(msg map[string]interface{})
acceptFn func(token string) bool
cmdPending map[string]chan map[string]interface{} // reqID -> 结果 channel
results map[string]resultEntry // reqID -> 已留档结果
}
// resultEntry 保存一次 cmdrun 的结果(供 device_ctl_cmdresult 查询)。
type resultEntry struct {
Result map[string]interface{}
Time time.Time
}
func NewRegistry() *Registry {
return &Registry{
devices: make(map[string]*DeviceMeta),
authorized: make(map[string]bool),
conns: make(map[string]*wconn),
onlineCh: make(chan string, 16),
cmdPending: make(map[string]chan map[string]interface{}),
results: make(map[string]resultEntry),
}
}
// SetAcceptToken 设置绑定 token 校验函数(插件注入,来自 Settings
func (r *Registry) SetAcceptToken(fn func(token string) bool) {
r.mu.Lock()
defer r.mu.Unlock()
r.acceptFn = fn
}
// SetStatusHandler 注册设备上报状态的回调。
func (r *Registry) SetStatusHandler(h func(msg map[string]interface{})) {
r.mu.Lock()
defer r.mu.Unlock()
r.onStatus = h
}
func (r *Registry) acceptBind(token string) bool {
r.mu.RLock()
fn := r.acceptFn
r.mu.RUnlock()
if fn == nil {
return false
}
return fn(token)
}
// ============ 设备查询 ============
// Online 返回设备是否在线。
func (r *Registry) Online(id string) bool {
r.mu.RLock()
defer r.mu.RUnlock()
m, ok := r.devices[id]
return ok && m.Online
}
// Authorized 返回设备是否已授权。
func (r *Registry) Authorized(id string) bool {
r.mu.RLock()
defer r.mu.RUnlock()
return r.authorized[id]
}
// List 返回全部设备(在线或历史),合并授权态。
func (r *Registry) List() []DeviceMeta {
r.mu.RLock()
defer r.mu.RUnlock()
out := make([]DeviceMeta, 0, len(r.devices))
for _, m := range r.devices {
c := *m
c.Authorized = r.authorized[c.DeviceID]
out = append(out, c)
}
return out
}
// OnlineList 返回在线的设备列表。
func (r *Registry) OnlineList() []DeviceMeta {
r.mu.RLock()
defer r.mu.RUnlock()
var out []DeviceMeta
for _, m := range r.devices {
if m.Online {
c := *m
c.Authorized = r.authorized[c.DeviceID]
out = append(out, c)
}
}
return out
}
// Get 返回单个设备。
func (r *Registry) Get(id string) (DeviceMeta, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
m, ok := r.devices[id]
if !ok {
return DeviceMeta{}, false
}
c := *m
c.Authorized = r.authorized[id]
return c, true
}
// ============ 授权 ============
// SetAuthorized 标记某设备已授权/取消授权(持久化由插件负责)。
func (r *Registry) SetAuthorized(id string, auth bool) {
r.mu.Lock()
defer r.mu.Unlock()
r.authorized[id] = auth
if m, ok := r.devices[id]; ok {
m.Authorized = auth
}
}
// RestoreAuthorized 插件启动时从配置恢复已授权设备集合。
func (r *Registry) RestoreAuthorized(ids []string) {
r.mu.Lock()
defer r.mu.Unlock()
for _, id := range ids {
r.authorized[id] = true
}
}
// AuthorizedIDs 返回全部已授权设备 ID供插件持久化
func (r *Registry) AuthorizedIDs() []string {
r.mu.RLock()
defer r.mu.RUnlock()
var out []string
for id, ok := range r.authorized {
if ok {
out = append(out, id)
}
}
return out
}
// ============ 在线状态维护 ============
func (r *Registry) register(meta DeviceMeta) {
r.mu.Lock()
meta.Online = true
meta.LastSeen = time.Now().Unix()
meta.Authorized = r.authorized[meta.DeviceID]
r.devices[meta.DeviceID] = &meta
r.mu.Unlock()
r.notifyChange(meta.DeviceID)
}
func (r *Registry) markOffline(id string) {
r.mu.Lock()
if m, ok := r.devices[id]; ok {
m.Online = false
}
delete(r.conns, id)
r.mu.Unlock()
r.notifyChange(id)
}
func (r *Registry) notifyChange(id string) {
select {
case r.onlineCh <- id:
default:
}
}
// ChangeChan 返回设备上下线变更通知。
func (r *Registry) ChangeChan() <-chan string { return r.onlineCh }
// SaveResult 保存一次命令执行结果(供 cmdresult 查询)。
func (r *Registry) SaveResult(reqID string, res map[string]interface{}) {
r.mu.Lock()
defer r.mu.Unlock()
if r.results == nil {
r.results = make(map[string]resultEntry)
}
r.results[reqID] = resultEntry{Result: res, Time: time.Now()}
}
// GetResult 返回某次命令执行的结果。
func (r *Registry) GetResult(reqID string) (map[string]interface{}, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
e, ok := r.results[reqID]
if !ok {
return nil, false
}
return e.Result, true
}
// ============ Pushagent -> 设备) ============
// PushJSON 向在线设备推送一条 JSON 消息。
func (r *Registry) PushJSON(deviceID string, payload map[string]interface{}) error {
r.mu.RLock()
c, ok := r.conns[deviceID]
r.mu.RUnlock()
if !ok {
return fmt.Errorf("device %s not online", deviceID)
}
return writeText(c.w, mustJSON(payload))
}
// PushCmd 向设备发送命令执行请求。
func (r *Registry) PushCmd(deviceID, reqID, command, cmdType string) error {
if cmdType == "" {
cmdType = "shell"
}
return r.PushJSON(deviceID, map[string]interface{}{
"op": "cmd",
"req_id": reqID,
"command": command,
"cmd_type": cmdType,
})
}
// PushData 向设备分块下发二进制数据(网关→设备,如 TTS 音频)。
// 协议(与 GUI 设备桥协商):
// 文本帧 cmd_speech_start {op, req_id, kind, mime, total} → N 个二进制帧(0x2, ≤8KB) → 文本帧 cmd_speech_end {op, req_id}
// kind 为语义标记(如 speechmime 为数据 MIME 类型。设备聚合后按自身能力处理(播放等)。
func (r *Registry) PushData(deviceID, reqID, kind, mime string, data []byte) error {
r.mu.RLock()
c, ok := r.conns[deviceID]
r.mu.RUnlock()
if !ok {
return fmt.Errorf("device %s not online", deviceID)
}
if err := writeText(c.w, mustJSON(map[string]interface{}{
"op": "cmd_speech_start",
"req_id": reqID,
"kind": kind,
"mime": mime,
"total": len(data),
})); err != nil {
return fmt.Errorf("push data start: %w", err)
}
const chunkSize = 8192
for off := 0; off < len(data); off += chunkSize {
end := off + chunkSize
if end > len(data) {
end = len(data)
}
if err := writeBinary(c.w, data[off:end]); err != nil {
return fmt.Errorf("push data chunk: %w", err)
}
}
if err := writeText(c.w, mustJSON(map[string]interface{}{
"op": "cmd_speech_end",
"req_id": reqID,
})); err != nil {
return fmt.Errorf("push data end: %w", err)
}
return nil
}
// AwaitResult 等待某请求的结果(带超时)。
func (r *Registry) AwaitResult(reqID string, timeout time.Duration) (map[string]interface{}, error) {
ch := make(chan map[string]interface{}, 1)
r.mu.Lock()
r.cmdPending[reqID] = ch
r.mu.Unlock()
defer func() {
r.mu.Lock()
delete(r.cmdPending, reqID)
r.mu.Unlock()
}()
select {
case res := <-ch:
return res, nil
case <-time.After(timeout):
return nil, fmt.Errorf("timeout waiting for device result")
}
}
// deliverResult 设备回执结果时由 handleWS 调用。
func (r *Registry) deliverResult(reqID string, res map[string]interface{}) {
r.mu.RLock()
ch, ok := r.cmdPending[reqID]
r.mu.RUnlock()
if ok {
select {
case ch <- res:
default:
}
}
}
// ============ WS 网关(标准库 HijackerRFC6455 子集) ============
const wsGUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
func wsAccept(key string) string {
h := sha256.Sum256([]byte(key + wsGUID))
return base64.StdEncoding.EncodeToString(h[:])
}
func headerListHas(h, sub string) bool {
for _, part := range strings.Split(h, ",") {
if strings.EqualFold(strings.TrimSpace(part), sub) {
return true
}
}
return false
}
func httpUpgrade(w http.ResponseWriter, r *http.Request) (net.Conn, *bufio.ReadWriter, error) {
if !headerListHas(r.Header.Get("Upgrade"), "websocket") ||
!headerListHas(r.Header.Get("Connection"), "upgrade") {
return nil, nil, fmt.Errorf("not a websocket upgrade request")
}
key := r.Header.Get("Sec-WebSocket-Key")
if key == "" {
return nil, nil, fmt.Errorf("missing Sec-WebSocket-Key")
}
hj, ok := w.(http.Hijacker)
if !ok {
return nil, nil, fmt.Errorf("hijack not supported")
}
conn, rw, err := hj.Hijack()
if err != nil {
return nil, nil, err
}
accept := wsAccept(key)
resp := "HTTP/1.1 101 Switching Protocols" + "\x0d\x0a" +
"Upgrade: websocket\x0d\x0aConnection: Upgrade\x0d\x0a" +
"Sec-WebSocket-Accept: " + accept + "\x0d\x0a\x0d\x0a"
if _, err := rw.WriteString(resp); err != nil {
conn.Close()
return nil, nil, err
}
if err := rw.Flush(); err != nil {
conn.Close()
return nil, nil, err
}
return conn, rw, nil
}
// readFrame 读取一个 WS 帧。返回 (payload, isClose, err)。
// opcode: 0x1 文本 / 0x2 二进制(设备→网关大体积数据分块,如录像回传)。
func readFrame(r *bufio.Reader) ([]byte, bool, byte, error) {
b0, err := r.ReadByte()
if err != nil {
return nil, true, 0, err
}
opcode := b0 & 0x0f
b1, err := r.ReadByte()
if err != nil {
return nil, true, opcode, err
}
masked := b1&0x80 != 0
length := uint64(b1 & 0x7f)
if length == 126 {
var ext [2]byte
if _, err := io.ReadFull(r, ext[:]); err != nil {
return nil, true, opcode, err
}
length = uint64(binary.BigEndian.Uint16(ext[:]))
} else if length == 127 {
var ext [8]byte
if _, err := io.ReadFull(r, ext[:]); err != nil {
return nil, true, opcode, err
}
length = binary.BigEndian.Uint64(ext[:])
}
// 二进制帧允许更大(录像分块聚合,单帧仍限 8MB 防滥用)
maxFrame := uint64(1 << 20)
if opcode == 0x2 {
maxFrame = 8 << 20
}
if length > maxFrame {
return nil, true, opcode, fmt.Errorf("frame too large")
}
var maskKey [4]byte
if masked {
if _, err := io.ReadFull(r, maskKey[:]); err != nil {
return nil, true, opcode, err
}
}
payload := make([]byte, length)
if _, err := io.ReadFull(r, payload); err != nil {
return nil, true, opcode, err
}
if masked {
for i := range payload {
payload[i] ^= maskKey[i%4]
}
}
switch opcode {
case 0x1, 0x2:
return payload, false, opcode, nil
case 0x8:
return nil, true, opcode, nil
case 0xa:
return nil, false, opcode, nil
case 0x9:
return nil, false, opcode, errPing
default:
return nil, false, opcode, fmt.Errorf("unsupported opcode %x", opcode)
}
}
var errPing = fmt.Errorf("ping")
func writeText(w *bufio.Writer, payload []byte) error {
return writeFrame(w, 0x1, payload)
}
// writeBinary 发送 WS 二进制帧0x2网关→设备大体积数据如 TTS 音频)分块下发。
func writeBinary(w *bufio.Writer, payload []byte) error {
return writeFrame(w, 0x2, payload)
}
func writeFrame(w *bufio.Writer, opcode byte, payload []byte) error {
if err := writeFrameHeader(w, opcode, len(payload)); err != nil {
return err
}
if _, err := w.Write(payload); err != nil {
return err
}
return w.Flush()
}
func writePong(w *bufio.Writer) error {
return writeFrameHeader(w, 0xa, 0)
}
func writeFrameHeader(w *bufio.Writer, opcode byte, length int) error {
if err := w.WriteByte(0x80 | opcode); err != nil {
return err
}
if length < 126 {
if err := w.WriteByte(byte(length)); err != nil {
return err
}
} else if length <= 0xffff {
if err := w.WriteByte(126); err != nil {
return err
}
var ext [2]byte
binary.BigEndian.PutUint16(ext[:], uint16(length))
if _, err := w.Write(ext[:]); err != nil {
return err
}
} else {
if err := w.WriteByte(127); err != nil {
return err
}
var ext [8]byte
binary.BigEndian.PutUint64(ext[:], uint64(length))
if _, err := w.Write(ext[:]); err != nil {
return err
}
}
return nil
}
func mustJSON(v interface{}) []byte {
b, err := json.Marshal(v)
if err != nil {
return []byte("{}")
}
return b
}
// ServeWS 是 WS 端点的 HTTP handler认证 tokenquery 或 Sec-WebSocket-Protocol
// 升级后进入 handleWS。未带 token 也允许 hello登记设备bind 时校验。
func (r *Registry) ServeWS(w http.ResponseWriter, req *http.Request) {
if req.Method != http.MethodGet {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
token := req.URL.Query().Get("token")
if token == "" {
for _, p := range req.Header.Values("Sec-WebSocket-Protocol") {
if strings.HasPrefix(p, "homeagent.") {
token = strings.TrimPrefix(p, "homeagent.")
break
}
}
}
if token != "" && !r.acceptBind(token) {
http.Error(w, "unauthorized", http.StatusUnauthorized)
return
}
conn, rw, err := httpUpgrade(w, req)
if err != nil {
http.Error(w, "upgrade failed: "+err.Error(), http.StatusBadRequest)
return
}
log.Printf("[remotedevice] ws connected from %s", conn.RemoteAddr())
go r.handleWS(conn, rw)
}
func (r *Registry) handleWS(conn net.Conn, rw *bufio.ReadWriter) {
defer conn.Close()
var curID string
defer func() {
if curID != "" {
r.markOffline(curID)
}
}()
// 二进制分块聚合状态(设备→网关,如录像回传):
// cmd_data_start 开启 → 0x2 帧追加 → cmd_data_end 聚合存入 cmdresult
var dataAccum *dataAccumulator
for {
payload, isClose, opcode, err := readFrame(rw.Reader)
if err != nil {
if err == errPing {
if werr := writePong(rw.Writer); werr != nil {
return
}
continue
}
return
}
if isClose {
return
}
if opcode == 0x2 {
// 二进制帧:处于聚合状态时追加数据块,否则忽略
if dataAccum != nil {
dataAccum.chunks = append(dataAccum.chunks, payload)
dataAccum.got += len(payload)
// 防滥用:超出声明 total 的 2 倍或硬上限 64MB 时放弃聚合
limit := int64(dataAccum.total)*2 + 1024
if limit < 64<<20 {
limit = 64 << 20
}
if int64(dataAccum.got) > limit {
log.Printf("[remotedevice] data accumulation exceeded limit for req %s, dropped", dataAccum.reqID)
dataAccum = nil
}
}
continue
}
var msg map[string]interface{}
if err := json.Unmarshal(payload, &msg); err != nil {
continue
}
op, _ := msg["op"].(string)
switch op {
case "hello":
meta := metaFromMsg(msg)
if meta.DeviceID == "" {
continue
}
meta.RemoteAddr = conn.RemoteAddr().String()
curID = meta.DeviceID
r.register(meta)
r.mu.Lock()
r.conns[meta.DeviceID] = &wconn{deviceID: meta.DeviceID, w: rw.Writer}
r.mu.Unlock()
if err := writeText(rw.Writer, mustJSON(map[string]interface{}{
"op": "hello_ack",
"device": meta.DeviceID,
"online": true,
})); err != nil {
return
}
case "bind":
token, _ := msg["token"].(string)
if r.acceptBind(token) {
id, _ := msg["device_id"].(string)
if id != "" {
// 默认不授权bind 仅验证 token + 登记设备;授权完全由用户手动
// GUI 设备页 / REST /api/v1/device/auth控制绝不自动授权。
}
if err := writeText(rw.Writer, mustJSON(map[string]interface{}{"op": "bind_ack", "ok": true})); err != nil {
return
}
} else {
if err := writeText(rw.Writer, mustJSON(map[string]interface{}{"op": "bind_ack", "ok": false, "error": "bad token"})); err != nil {
return
}
}
case "status":
id, _ := msg["device_id"].(string)
r.mu.Lock()
if m, ok := r.devices[id]; ok {
m.LastSeen = time.Now().Unix()
}
r.mu.Unlock()
r.mu.RLock()
h := r.onStatus
r.mu.RUnlock()
if h != nil {
h(msg)
}
case "cmd_result":
reqID, _ := msg["req_id"].(string)
if reqID != "" {
r.deliverResult(reqID, msg)
}
case "cmd_data_start":
reqID, _ := msg["req_id"].(string)
if reqID == "" {
continue
}
total, _ := msg["total"].(float64)
kind, _ := msg["kind"].(string)
mime, _ := msg["mime"].(string)
dataAccum = &dataAccumulator{
reqID: reqID,
kind: kind,
mime: mime,
total: int(total),
}
case "cmd_data_end":
reqID, _ := msg["req_id"].(string)
status, _ := msg["status"].(string)
if dataAccum == nil || dataAccum.reqID != reqID {
continue
}
acc := dataAccum
dataAccum = nil
if status != "ok" {
r.SaveResult(reqID, map[string]interface{}{
"op": "cmd_result", "req_id": reqID, "status": "error",
"error": "device reported transfer failure",
})
r.deliverResult(reqID, map[string]interface{}{
"op": "cmd_result", "req_id": reqID, "status": "error",
"error": "device reported transfer failure",
})
continue
}
data := make([]byte, 0, acc.got)
for _, c := range acc.chunks {
data = append(data, c...)
}
res := map[string]interface{}{
"op": "cmd_result",
"req_id": reqID,
"status": "ok",
"kind": acc.kind,
"mime": acc.mime,
"size": len(data),
"expected": acc.total,
// base64 编码完整二进制(录像 mp4 等),供 agent/上层取回后解码使用
"data_base64": base64.StdEncoding.EncodeToString(data),
}
r.SaveResult(reqID, res)
r.deliverResult(reqID, res)
}
}
}
// dataAccumulator 聚合设备→网关的二进制分块传输(如录像回传)。
type dataAccumulator struct {
reqID string
kind string
mime string
total int
chunks [][]byte
got int
}
func metaFromMsg(msg map[string]interface{}) DeviceMeta {
var meta DeviceMeta
if d, ok := msg["device"].(map[string]interface{}); ok {
if v, ok := d["device_id"].(string); ok {
meta.DeviceID = v
}
if v, ok := d["name"].(string); ok {
meta.Name = v
}
if v, ok := d["kind"].(string); ok {
meta.Kind = v
}
if caps, ok := d["caps"].([]interface{}); ok {
for _, c := range caps {
if s, ok := c.(string); ok {
meta.Caps = append(meta.Caps, s)
}
}
}
if info, ok := d["info"].(map[string]interface{}); ok {
if len(info) > 0 {
meta.Info = info
}
}
}
return meta
}