2026-06-01 14:31:15 +08:00

183 lines
3.7 KiB
Go

package main
import (
"encoding/binary"
"encoding/json"
"errors"
"flag"
"fmt"
"io"
"net"
"os"
"time"
)
const (
defaultSocketPath = "/var/lib/iot/ctl.sock"
maxPacketSize = 65535
dialTimeout = 5 * time.Second
requestTimeout = 10 * time.Second
)
type request struct {
Method string `json:"method"`
Params map[string]string `json:"params,omitempty"`
}
type response struct {
Result json.RawMessage `json:"result,omitempty"`
Error *responseError `json:"error,omitempty"`
}
type responseError struct {
Message string `json:"message"`
}
func main() {
if err := run(os.Args[1:]); err != nil {
fmt.Fprintf(os.Stderr, "iot_ctrl: %v\n", err)
os.Exit(1)
}
}
func run(args []string) error {
flags := flag.NewFlagSet("iot_ctrl", flag.ContinueOnError)
flags.SetOutput(io.Discard)
socketPath := defaultSocketPath
if envSocketPath := os.Getenv("IOT_CTRL_SOCKET_PATH"); envSocketPath != "" {
socketPath = envSocketPath
}
flags.StringVar(&socketPath, "socket", socketPath, "control socket path")
if err := flags.Parse(args); err != nil {
return usageError()
}
rest := flags.Args()
if len(rest) == 0 {
return usageError()
}
req, err := parseCommand(rest)
if err != nil {
return err
}
result, err := call(socketPath, req)
if err != nil {
return err
}
if result != "" {
fmt.Println(result)
}
return nil
}
func parseCommand(args []string) (request, error) {
switch args[0] {
case "ping":
if len(args) != 1 {
return request{}, usageError()
}
return request{Method: "ping"}, nil
case "add-client":
flags := flag.NewFlagSet("add-client", flag.ContinueOnError)
flags.SetOutput(io.Discard)
var uuid string
var token string
flags.StringVar(&uuid, "uuid", "", "client uuid")
flags.StringVar(&token, "token", "", "client token")
if err := flags.Parse(args[1:]); err != nil {
return request{}, usageError()
}
if uuid == "" || token == "" || flags.NArg() != 0 {
return request{}, usageError()
}
return request{
Method: "add_client",
Params: map[string]string{
"uuid": uuid,
"token": token,
},
}, nil
default:
return request{}, usageError()
}
}
func call(socketPath string, req request) (string, error) {
conn, err := net.DialTimeout("unix", socketPath, dialTimeout)
if err != nil {
return "", err
}
defer conn.Close()
if err := conn.SetDeadline(time.Now().Add(requestTimeout)); err != nil {
return "", err
}
body, err := json.Marshal(req)
if err != nil {
return "", err
}
if err := writePacket(conn, body); err != nil {
return "", err
}
reply, err := readPacket(conn)
if err != nil {
return "", err
}
var resp response
if err := json.Unmarshal(reply, &resp); err != nil {
return "", err
}
if resp.Error != nil {
return "", errors.New(resp.Error.Message)
}
if len(resp.Result) == 0 {
return "", nil
}
var text string
if err := json.Unmarshal(resp.Result, &text); err == nil {
return text, nil
}
return string(resp.Result), nil
}
func writePacket(w io.Writer, body []byte) error {
if len(body) > maxPacketSize {
return fmt.Errorf("request exceeds packet=2 limit: %d bytes", len(body))
}
header := make([]byte, 2)
binary.BigEndian.PutUint16(header, uint16(len(body)))
if _, err := w.Write(header); err != nil {
return err
}
_, err := w.Write(body)
return err
}
func readPacket(r io.Reader) ([]byte, error) {
header := make([]byte, 2)
if _, err := io.ReadFull(r, header); err != nil {
return nil, err
}
size := binary.BigEndian.Uint16(header)
body := make([]byte, int(size))
if _, err := io.ReadFull(r, body); err != nil {
return nil, err
}
return body, nil
}
func usageError() error {
return errors.New("usage: iot_ctrl [--socket path] ping | iot_ctrl [--socket path] add-client --uuid <uuid> --token <token>")
}