mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
341 lines
7.8 KiB
Go
341 lines
7.8 KiB
Go
package main
|
|
|
|
import (
|
|
"bufio"
|
|
"context"
|
|
"crypto/rand"
|
|
"encoding/base64"
|
|
"encoding/binary"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"os"
|
|
"os/exec"
|
|
"regexp"
|
|
"runtime"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
type deviceBridge struct {
|
|
addr string
|
|
token string
|
|
deviceID string
|
|
conn net.Conn
|
|
mu sync.Mutex
|
|
stop chan struct{}
|
|
}
|
|
|
|
func newDeviceBridge(addr, token string) *deviceBridge {
|
|
return &deviceBridge{addr: addr, token: token, stop: make(chan struct{})}
|
|
}
|
|
|
|
func (b *deviceBridge) Start() error {
|
|
if b.addr == "" || b.token == "" {
|
|
return fmt.Errorf("device bridge: addr/token required")
|
|
}
|
|
host, port := b.addr, "9890"
|
|
if h, p, err := net.SplitHostPort(b.addr); err == nil {
|
|
host, port = h, p
|
|
} else if idx := strings.LastIndex(b.addr, ":"); idx >= 0 {
|
|
host = strings.TrimPrefix(b.addr[:idx], "http://")
|
|
port = b.addr[idx+1:]
|
|
}
|
|
hostname, _ := os.Hostname()
|
|
if hostname == "" {
|
|
hostname = "local"
|
|
}
|
|
b.deviceID = "waiter-" + sanitizeID(hostname)
|
|
|
|
conn, err := net.DialTimeout("tcp", net.JoinHostPort(host, port), 5*time.Second)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
b.conn = conn
|
|
|
|
key := wsKey()
|
|
path := "/api/v1/device/ws?token=" + urlEscape(b.token)
|
|
var sb strings.Builder
|
|
sb.WriteString("GET " + path + " HTTP/1.1" + CRLF)
|
|
sb.WriteString("Host: " + host + ":" + port + CRLF)
|
|
sb.WriteString("Upgrade: websocket" + CRLF)
|
|
sb.WriteString("Connection: Upgrade" + CRLF)
|
|
sb.WriteString("Sec-WebSocket-Key: " + key + CRLF)
|
|
sb.WriteString("Sec-WebSocket-Version: 13" + CRLF + CRLF)
|
|
if _, err := conn.Write([]byte(sb.String())); err != nil {
|
|
conn.Close()
|
|
return err
|
|
}
|
|
br := bufio.NewReader(conn)
|
|
headerBuf := ""
|
|
for {
|
|
line, err := br.ReadString('\n')
|
|
if err != nil {
|
|
conn.Close()
|
|
return err
|
|
}
|
|
headerBuf += line
|
|
if strings.Contains(headerBuf, CRLF+CRLF) {
|
|
break
|
|
}
|
|
}
|
|
if !strings.Contains(headerBuf, " 101 ") {
|
|
conn.Close()
|
|
return fmt.Errorf("ws upgrade failed: %s", firstLine(headerBuf))
|
|
}
|
|
|
|
b.sendJSON(helloMsg(b.deviceID))
|
|
b.sendJSON(map[string]interface{}{"op": "bind", "device_id": b.deviceID, "token": b.token})
|
|
|
|
go b.readLoop(br)
|
|
return nil
|
|
}
|
|
|
|
// CRLF 用常量避免转义地狱
|
|
const CRLF = "\r\n"
|
|
|
|
func helloMsg(deviceID string) map[string]interface{} {
|
|
hostname, _ := os.Hostname()
|
|
return map[string]interface{}{
|
|
"op": "hello",
|
|
"device": map[string]interface{}{
|
|
"device_id": deviceID,
|
|
"name": "HomeAgent CLI (waiter)",
|
|
"kind": "computer",
|
|
"caps": []string{"status", "cmdrun", "deviceinfo"},
|
|
"info": map[string]interface{}{
|
|
"hostname": hostname,
|
|
"platform": runtime.GOOS,
|
|
"arch": runtime.GOARCH,
|
|
"cpus": runtime.NumCPU(),
|
|
"mem_mb": memTotalMB(),
|
|
},
|
|
},
|
|
}
|
|
}
|
|
|
|
func memTotalMB() int64 {
|
|
if runtime.GOOS != "linux" {
|
|
return 0
|
|
}
|
|
data, err := os.ReadFile("/proc/meminfo")
|
|
if err != nil {
|
|
return 0
|
|
}
|
|
for _, line := range strings.Split(string(data), "\n") {
|
|
if strings.HasPrefix(line, "MemTotal:") {
|
|
f := strings.Fields(line)
|
|
if len(f) >= 2 {
|
|
var kb int64
|
|
fmt.Sscanf(f[1], "%d", &kb)
|
|
return kb / 1024
|
|
}
|
|
}
|
|
}
|
|
return 0
|
|
}
|
|
|
|
func (b *deviceBridge) sendJSON(v interface{}) {
|
|
payload, err := json.Marshal(v)
|
|
if err != nil {
|
|
return
|
|
}
|
|
b.mu.Lock()
|
|
defer b.mu.Unlock()
|
|
if b.conn == nil {
|
|
return
|
|
}
|
|
writeWSFrame(b.conn, 0x1, payload)
|
|
}
|
|
|
|
func writeWSFrame(conn net.Conn, opcode byte, payload []byte) {
|
|
length := len(payload)
|
|
hdr := []byte{0x80 | opcode}
|
|
switch {
|
|
case length < 126:
|
|
hdr = append(hdr, 0x80|byte(length))
|
|
case length <= 0xffff:
|
|
hdr = append(hdr, 0x80|126, 0, 0)
|
|
binary.BigEndian.PutUint16(hdr[len(hdr)-2:], uint16(length))
|
|
default:
|
|
hdr = append(hdr, 0x80|127)
|
|
hdr = append(hdr, 0, 0, 0, 0, 0, 0, 0, 0)
|
|
binary.BigEndian.PutUint64(hdr[len(hdr)-8:], uint64(length))
|
|
}
|
|
var mask [4]byte
|
|
rand.Read(mask[:])
|
|
masked := make([]byte, length)
|
|
for i := 0; i < length; i++ {
|
|
masked[i] = payload[i] ^ mask[i&3]
|
|
}
|
|
conn.Write(append(append(hdr, mask[:]...), masked...))
|
|
}
|
|
|
|
func (b *deviceBridge) readLoop(br *bufio.Reader) {
|
|
for {
|
|
select {
|
|
case <-b.stop:
|
|
return
|
|
default:
|
|
}
|
|
payload, err := readWSFrame(br)
|
|
if err != nil {
|
|
return
|
|
}
|
|
var msg map[string]interface{}
|
|
if err := json.Unmarshal(payload, &msg); err != nil {
|
|
continue
|
|
}
|
|
if msg["op"] != "cmd" {
|
|
continue
|
|
}
|
|
reqID, _ := msg["req_id"].(string)
|
|
command, _ := msg["command"].(string)
|
|
if reqID == "" {
|
|
continue
|
|
}
|
|
go b.runCommand(reqID, command)
|
|
}
|
|
}
|
|
|
|
func readWSFrame(br *bufio.Reader) ([]byte, error) {
|
|
b0, err := br.ReadByte()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
opcode := b0 & 0x0f
|
|
b1, err := br.ReadByte()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
length := uint64(b1 & 0x7f)
|
|
if length == 126 {
|
|
var ext [2]byte
|
|
if _, err := io.ReadFull(br, ext[:]); err != nil {
|
|
return nil, err
|
|
}
|
|
length = uint64(binary.BigEndian.Uint16(ext[:]))
|
|
} else if length == 127 {
|
|
var ext [8]byte
|
|
if _, err := io.ReadFull(br, ext[:]); err != nil {
|
|
return nil, err
|
|
}
|
|
length = binary.BigEndian.Uint64(ext[:])
|
|
}
|
|
if length > 1<<20 {
|
|
return nil, fmt.Errorf("frame too large")
|
|
}
|
|
payload := make([]byte, length)
|
|
if _, err := io.ReadFull(br, payload); err != nil {
|
|
return nil, err
|
|
}
|
|
switch opcode {
|
|
case 0x1:
|
|
return payload, nil
|
|
case 0x8:
|
|
return nil, io.EOF
|
|
default:
|
|
return nil, nil
|
|
}
|
|
}
|
|
|
|
var waiterAllowCmd = regexp.MustCompile("^(ls|pwd|whoami|uname|date|echo|uptime|hostname|cat|df|free|ps|ip|dir|node|python3?|npm|git|curl|wget|systeminfo|tasklist)\\b")
|
|
|
|
func (b *deviceBridge) runCommand(reqID, command string) {
|
|
res := b.execCommand(command)
|
|
res["op"] = "cmd_result"
|
|
res["req_id"] = reqID
|
|
res["device_id"] = b.deviceID
|
|
b.sendJSON(res)
|
|
}
|
|
|
|
func (b *deviceBridge) execCommand(cmdLine string) map[string]interface{} {
|
|
if strings.TrimSpace(cmdLine) == "" {
|
|
return map[string]interface{}{"status": "error", "output": "empty command", "exit_code": -1}
|
|
}
|
|
if !waiterAllowCmd.MatchString(strings.TrimSpace(cmdLine)) {
|
|
return map[string]interface{}{"status": "error", "output": "command not in whitelist", "exit_code": -1}
|
|
}
|
|
parts := strings.Fields(strings.TrimSpace(cmdLine))
|
|
if len(parts) == 0 {
|
|
return map[string]interface{}{"status": "error", "output": "empty", "exit_code": -1}
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
|
defer cancel()
|
|
cmd := exec.CommandContext(ctx, parts[0], parts[1:]...)
|
|
out, err := cmd.Output()
|
|
if err != nil {
|
|
if ctx.Err() != nil {
|
|
return map[string]interface{}{"status": "error", "output": "timeout", "exit_code": -2}
|
|
}
|
|
return map[string]interface{}{"status": "error", "output": truncate8k(string(out)) + " err: " + err.Error(), "exit_code": -1}
|
|
}
|
|
return map[string]interface{}{"status": "ok", "output": truncate8k(string(out)), "exit_code": 0}
|
|
}
|
|
|
|
const devMaxOut = 8192
|
|
|
|
func truncate8k(s string) string {
|
|
if len(s) <= devMaxOut {
|
|
return s
|
|
}
|
|
return s[:devMaxOut]
|
|
}
|
|
|
|
func firstLine(s string) string {
|
|
if i := strings.IndexByte(s, '\n'); i >= 0 {
|
|
return strings.TrimSpace(s[:i])
|
|
}
|
|
return strings.TrimSpace(s)
|
|
}
|
|
|
|
func sanitizeID(s string) string {
|
|
var sb strings.Builder
|
|
for _, r := range s {
|
|
if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '-' || r == '_' {
|
|
sb.WriteRune(r)
|
|
} else {
|
|
sb.WriteByte('_')
|
|
}
|
|
}
|
|
return sb.String()
|
|
}
|
|
|
|
func urlEscape(s string) string {
|
|
var sb strings.Builder
|
|
const hex = "0123456789ABCDEF"
|
|
for i := 0; i < len(s); i++ {
|
|
c := s[i]
|
|
if (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '-' || c == '_' || c == '.' || c == '~' {
|
|
sb.WriteByte(c)
|
|
} else {
|
|
sb.WriteByte('%')
|
|
sb.WriteByte(hex[c>>4])
|
|
sb.WriteByte(hex[c&0xf])
|
|
}
|
|
}
|
|
return sb.String()
|
|
}
|
|
|
|
func wsKey() string {
|
|
var b [16]byte
|
|
rand.Read(b[:])
|
|
return base64.StdEncoding.EncodeToString(b[:])
|
|
}
|
|
|
|
func (b *deviceBridge) Stop() {
|
|
b.mu.Lock()
|
|
defer b.mu.Unlock()
|
|
select {
|
|
case <-b.stop:
|
|
default:
|
|
close(b.stop)
|
|
}
|
|
if b.conn != nil {
|
|
b.conn.Close()
|
|
b.conn = nil
|
|
}
|
|
}
|