Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b3dcf6ed66 | ||
|
|
512b0a11e6 | ||
|
|
41079fa1c8 | ||
|
|
099ae7468e | ||
|
|
2e875c9acc | ||
|
|
601a0fa450 | ||
|
|
1ae612cb6f | ||
|
|
f2998ea538 | ||
|
|
12c64dff17 |
@@ -0,0 +1,5 @@
|
|||||||
|
*.pcap
|
||||||
|
docs/
|
||||||
|
.env
|
||||||
|
.git/
|
||||||
|
/n95bridge
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
.env
|
||||||
|
/n95bridge
|
||||||
+28
@@ -0,0 +1,28 @@
|
|||||||
|
# syntax=docker/dockerfile:1
|
||||||
|
|
||||||
|
# Stage 1: Build the static Go binary
|
||||||
|
FROM golang:1.27.1@sha256:d8c54b23d9b0dcbf6e1e69bbfb9f8f2b73059bb29c54e1de2f495ec9dfcb8ef0 AS builder
|
||||||
|
|
||||||
|
WORKDIR /src
|
||||||
|
|
||||||
|
COPY go.mod go.sum ./
|
||||||
|
RUN go mod download
|
||||||
|
|
||||||
|
COPY cmd/ cmd/
|
||||||
|
COPY internal/ internal/
|
||||||
|
|
||||||
|
ARG TARGETOS
|
||||||
|
ARG TARGETARCH
|
||||||
|
|
||||||
|
RUN CGO_ENABLED=0 GOOS=${TARGETOS} GOARCH=${TARGETARCH} go build -trimpath -ldflags="-s -w" -o /n95bridge ./cmd/n95bridge
|
||||||
|
|
||||||
|
# Stage 2: Minimal non-root distroless runtime
|
||||||
|
FROM gcr.io/distroless/static-debian12:nonroot@sha256:6ec646aa5008ab558e658ec354ee11fa6770559e8697b0a88094ae4a899dd3b3
|
||||||
|
|
||||||
|
COPY --from=builder /n95bridge /n95bridge
|
||||||
|
|
||||||
|
USER 65532:65532
|
||||||
|
|
||||||
|
EXPOSE 8005 8007 5223 8080
|
||||||
|
|
||||||
|
ENTRYPOINT ["/n95bridge"]
|
||||||
@@ -0,0 +1,176 @@
|
|||||||
|
# Deebot N95 local MQTT bridge
|
||||||
|
|
||||||
|
This service redirects an **already provisioned Ecovacs Deebot N95** from its legacy cloud bootstrap and XMPP endpoint to a local bridge. It exposes the robot to Home Assistant through MQTT discovery as a native MQTT vacuum. The bridge supports multiple robots, with one MQTT client and one topic tree per robot serial number. It does not provision a robot or run a DNS server.
|
||||||
|
|
||||||
|
```text
|
||||||
|
N95 ── DNS ──> your LAN resolver
|
||||||
|
N95 ── HTTP 8007 / 8005, XMPP 5223 ──> n95bridge ──> MQTT broker ──> Home Assistant
|
||||||
|
```
|
||||||
|
|
||||||
|
The bridge's robot-facing XMPP connection is **plaintext**, including the robot's SASL PLAIN login. Run it on a trusted LAN and limit access to its listener ports. MQTT can be protected separately with TLS and broker authentication. The robot's existing password is accepted during the local XMPP handshake; there is no robot credential to enter into the bridge configuration.
|
||||||
|
|
||||||
|
## Before you start
|
||||||
|
|
||||||
|
You need:
|
||||||
|
|
||||||
|
- An already provisioned Deebot N95 using the `wukong` / class `155` protocol captured for this project. Other Ecovacs models are not established as compatible.
|
||||||
|
- A Docker host (or a machine running the Go binary) with a **stable LAN IPv4 address** reachable from the robot. Reserve its address in DHCP.
|
||||||
|
- Control of the DNS resolver supplied to the robot by DHCP. It must answer `lbo.ecouser.net` with that LAN address. A direct A record is sufficient. The bridge does not provide DNS. The XMPP JID domain `155.ecorobot.net` does not need a DNS override for this captured firmware.
|
||||||
|
- TCP **8007**, **8005**, and **5223** reachable from the robot at that address. TCP **8080** is the optional health endpoint; it need not be open to the robot.
|
||||||
|
- An MQTT broker reachable from the bridge and Home Assistant configured to use that broker, with MQTT discovery enabled. Its discovery prefix must match `HA_DISCOVERY_PREFIX` (normally `homeassistant`).
|
||||||
|
- A correct host clock and local timezone. The bridge sends the robot the current time and the process's local UTC offset on each session.
|
||||||
|
|
||||||
|
The N95 caches the XMPP IP and port after boot. If the bridge address changes while the robot is powered on, update DNS and **reboot the robot** so it performs bootstrap again. A normal XMPP reconnect goes directly to the cached address.
|
||||||
|
|
||||||
|
## Docker Compose setup
|
||||||
|
|
||||||
|
The supplied [docker-compose.yml](docker-compose.yml) uses published ports on Docker's default bridge network. Its `192.0.2.10` and `mqtt.example.invalid` values are **documentation placeholders**. Edit the file before starting it:
|
||||||
|
|
||||||
|
1. Reserve the Docker host's LAN IPv4 address. Set `ADVERTISE_IP` to that address and configure your LAN DNS resolver to return it for `lbo.ecouser.net`. Ensure the robot receives that resolver through DHCP.
|
||||||
|
2. Set `MQTT_HOST` to the broker's hostname or IP address **as seen from inside the container**. Enter a bare host, with no `tcp://`, `mqtt://`, or `:port`. If the broker runs on the Docker host, `localhost` inside the container refers to the container, so use an address or hostname the container can reach.
|
||||||
|
3. Set `TZ` to your local IANA timezone, for example `Europe/London`. Keep `BIND_ADDRESS: "0.0.0.0"` inside the container. `ADVERTISE_IP` is the host's LAN address, not this bind address or a Docker subnet address.
|
||||||
|
4. If your broker requires authentication, uncomment the `env_file` block and create an uncommitted `.env` beside the Compose file containing `MQTT_USERNAME=...` and `MQTT_PASSWORD=...`. The file is gitignored and excluded from the image build. Compose does **not** pass values from its automatic `.env` interpolation file into the container unless they are referenced in `environment` or included through `env_file`. Keep credentials out of the Compose file and source control.
|
||||||
|
5. If your broker uses TLS, set `MQTT_TLS: "true"` (default port 8883), and set `MQTT_PORT` in the Compose `environment` block if your broker uses a different port. For a private CA, mount a readable PEM certificate bundle into the container and set `MQTT_CA_FILE` to its **container path**. For example, add `volumes: ["./broker-ca.pem:/certs/broker-ca.pem:ro"]` and `MQTT_CA_FILE: "/certs/broker-ca.pem"`. The runtime is non-root (UID 65532), so the mounted file must be readable by that user. The CA is added to the system trust pool; TLS hostname verification remains enabled.
|
||||||
|
6. If you change a robot-facing port, change its `PORT_*` variable **and** the matching Compose published port. The host port must equal the port returned by `/lookup.do`. Leave the four listener ports distinct.
|
||||||
|
|
||||||
|
Start and inspect the service:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
docker compose up -d --build
|
||||||
|
docker compose ps
|
||||||
|
docker compose logs -f n95bridge
|
||||||
|
```
|
||||||
|
|
||||||
|
The image builds a static Go binary and runs it as a non-root user. The Compose health check executes `/n95bridge healthcheck` inside the container. On `SIGTERM`, the bridge attempts to publish retained `offline` before closing XMPP; the shutdown budget is five seconds.
|
||||||
|
|
||||||
|
The four PCAP files at the repository root are test fixtures containing captured traffic. They are not needed at runtime, and `.dockerignore` excludes them from the image. Treat those files as sensitive.
|
||||||
|
|
||||||
|
### Run without Docker
|
||||||
|
|
||||||
|
Install the Go version specified by [go.mod](go.mod), currently **Go 1.27.1**, and build from the repository root:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
go build -o n95bridge ./cmd/n95bridge
|
||||||
|
ADVERTISE_IP=192.0.2.10 MQTT_HOST=broker.example.net TZ=Europe/London ./n95bridge
|
||||||
|
```
|
||||||
|
|
||||||
|
Replace both example addresses. Set any other variables from the table below in the service manager or process environment. The binary reads environment variables; it does not parse `.env` itself. Arrange for it to start on boot, receive `SIGTERM` for orderly shutdown, and bind the configured ports. `./n95bridge healthcheck` uses the same environment and checks `http://127.0.0.1:${HEALTH_PORT:-8080}/healthz`.
|
||||||
|
|
||||||
|
## Configuration reference
|
||||||
|
|
||||||
|
Only `ADVERTISE_IP` and `MQTT_HOST` are required. An absent variable uses its default. Empty values are allowed for optional `MQTT_CA_FILE`, `MQTT_USERNAME`, and `MQTT_PASSWORD`; empty `MQTT_PORT` and `MQTT_CLIENT_ID` use their defaults. Other configured values must be nonempty. Boolean values must be exactly `true` or `false` in lowercase. All four listener ports must be different and within 1–65535.
|
||||||
|
|
||||||
|
| Variable | Default | Meaning and constraints |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| `ADVERTISE_IP` | required | Literal nonzero IPv4 returned in both lookup responses; the robot must be able to reach it for its powered-on lifetime. |
|
||||||
|
| `BIND_ADDRESS` | `0.0.0.0` | IP address on which all four listeners bind. It is never advertised. For the supplied container networking, keep `0.0.0.0`. |
|
||||||
|
| `PORT_LOOKUP` | `8007` | HTTP `POST /lookup.do` for `EcoMsgNew` and `EcoUpdate`. |
|
||||||
|
| `PORT_FIRMWARE` | `8005` | HTTP firmware check; returns the expected 404. |
|
||||||
|
| `PORT_XMPP` | `5223` | Plaintext robot XMPP listener; advertised for `EcoMsgNew`. |
|
||||||
|
| `HEALTH_PORT` | `8080` | HTTP `GET /healthz` and the `healthcheck` subcommand. |
|
||||||
|
| `MQTT_HOST` | required | Broker DNS name or address, without a scheme or port. Also used as the TLS server name. |
|
||||||
|
| `MQTT_PORT` | `1883` or `8883` | Broker port; defaults to 8883 when `MQTT_TLS=true`, otherwise 1883. |
|
||||||
|
| `MQTT_TLS` | `false` | Enable TLS from the start of the MQTT TCP connection. Uses the system CA pool and verifies the broker hostname. |
|
||||||
|
| `MQTT_CA_FILE` | unset | Optional readable PEM certificate bundle added to the TLS trust pool. Mount it into a container if applicable. |
|
||||||
|
| `MQTT_USERNAME` | unset | Optional broker username. |
|
||||||
|
| `MQTT_PASSWORD` | unset | Optional broker password; a nonempty password requires a username. The bridge omits the password from its configuration log. |
|
||||||
|
| `MQTT_CLIENT_ID` | `n95bridge-<hostname>` | Prefix for per-robot broker client IDs. Explicit values must match `[A-Za-z0-9_-]{1,64}`. Use distinct prefixes if running separate bridge instances against the same broker. |
|
||||||
|
| `MQTT_BASE` | `ecovacs` | Single MQTT topic level before `/<serial>`; letters, digits, `_`, and `-` only. |
|
||||||
|
| `HA_DISCOVERY_PREFIX` | `homeassistant` | Single topic level for Home Assistant MQTT discovery; letters, digits, `_`, and `-` only. Must match Home Assistant's setting. |
|
||||||
|
| `CONTROLLER_JID` | `n95bridge@ecouser.net/homeassistant` | Virtual XMPP sender JID in `local@domain/resource` form. The robot learns this address from bridge requests. |
|
||||||
|
| `RAW_COMMANDS` | `false` | Enable the optional raw `<ctl>` command input. Keep disabled unless you explicitly need protocol experimentation. |
|
||||||
|
| `LOG_LEVEL` | `info` | `debug`, `info`, `warn`, or `error` on stderr. |
|
||||||
|
|
||||||
|
`TZ` is a standard process timezone setting, not a bridge-specific option. It controls the local offset sent by `SetTime`; the binary includes timezone data for the distroless image. The default MQTT client ID prefix is derived from the hostname. The bridge appends each robot's serial to make its final broker client ID.
|
||||||
|
|
||||||
|
The bridge creates an MQTT client **only after a robot reaches its XMPP READY state and reveals its serial**. No MQTT connection at initial process startup is expected. It reconnects to the broker automatically when needed.
|
||||||
|
|
||||||
|
## Check the installation
|
||||||
|
|
||||||
|
From a machine on the robot's LAN, first verify the DNS response using the resolver the robot receives, then check the listeners. Replace the example addresses with your deployment values:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
nslookup lbo.ecouser.net 192.0.2.53
|
||||||
|
curl -sS -X POST -H 'Content-Type: application/json' \
|
||||||
|
--data '{"todo":"FindBest","service":"EcoMsgNew"}' \
|
||||||
|
http://192.0.2.10:8007/lookup.do
|
||||||
|
curl -sS -X POST -H 'Content-Type: application/json' \
|
||||||
|
--data '{"todo":"FindBest","service":"EcoUpdate"}' \
|
||||||
|
http://192.0.2.10:8007/lookup.do
|
||||||
|
curl -i http://192.0.2.10:8005/products/wukong/class/155/firmware/latest.json
|
||||||
|
curl -sS http://192.0.2.10:8080/healthz
|
||||||
|
```
|
||||||
|
|
||||||
|
The two lookup responses should contain the advertised LAN IP with numeric ports `5223` and `8005`, respectively. The firmware request should return **404** with body `Not Found`; `/healthz` should return `ok`. If the health port is bound only locally or blocked from your test machine, run the health check inside the container instead:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
docker compose exec n95bridge /n95bridge healthcheck
|
||||||
|
```
|
||||||
|
|
||||||
|
Boot or reboot the robot after the DNS override is active. The log line `sasl authenticated` includes `authcid=<serial>`, which identifies the robot's MQTT topic tree. Once its announce ping succeeds, `{MQTT_BASE}/{serial}/availability` becomes retained `online` and Home Assistant should discover a `Deebot N95` vacuum. A robot disconnect sets retained `offline`; an abrupt bridge failure is also covered by the broker's retained MQTT last will. Initial status, schedules, battery, and consumable requests follow the session announcement. State and attributes are retained so Home Assistant can recover them after a restart.
|
||||||
|
|
||||||
|
`/healthz` checks **process liveness only**. It can return `ok` while the robot is offline or the broker is unavailable. Use MQTT availability and bridge logs to diagnose connectivity.
|
||||||
|
|
||||||
|
## MQTT and Home Assistant interface
|
||||||
|
|
||||||
|
The serial is learned from the robot's XMPP login; it is **not** an environment variable. With defaults, the discovery topic is `homeassistant/vacuum/ecovacs_<serial>/config`, and the per-robot root is `ecovacs/<serial>`. Changing `MQTT_BASE` changes the per-robot topics, while the discovery object ID and unique ID remain `ecovacs_<serial>`.
|
||||||
|
|
||||||
|
| Topic under `{MQTT_BASE}/{serial}/` | Direction | Contents | Retained |
|
||||||
|
| --- | --- | --- | --- |
|
||||||
|
| `command` | to bridge | Exact string: `start`, `stop`, `return_to_base`, `clean_spot`, or `locate` | No |
|
||||||
|
| `set_fan_speed` | to bridge | Exact string: `standard` or `strong` | No |
|
||||||
|
| `send_command` | to bridge | JSON extension command described below | No |
|
||||||
|
| `availability` | from bridge | `online` or `offline` | Yes |
|
||||||
|
| `state` | from bridge | JSON with `state` and `fan_speed` | Yes |
|
||||||
|
| `json_attributes` | from bridge | Complete JSON snapshot of battery, consumables, schedules, and errors | Yes |
|
||||||
|
| `command_result` | from bridge | JSON command trace with `sid`, `cid`, `command`, `phase`, `ret`, `errno`, and `timestamp` | No |
|
||||||
|
| `raw` | from bridge | JSON for unparsed protocol data, including direction, timestamp, XML, and reason | No |
|
||||||
|
|
||||||
|
Discovery, availability, state, and attributes are published at MQTT QoS 0. The bridge subscribes to command topics at QoS 1. Do not publish retained command messages: an old command could be delivered again after a subscription or reconnect. The bridge republishes discovery when it reconnects to MQTT and when Home Assistant sends `online` to `{HA_DISCOVERY_PREFIX}/status`.
|
||||||
|
|
||||||
|
If your broker uses topic ACLs, allow each bridge client to **read** `{MQTT_BASE}/+/command`, `{MQTT_BASE}/+/set_fan_speed`, `{MQTT_BASE}/+/send_command`, and `{HA_DISCOVERY_PREFIX}/status`; allow it to **write** `{MQTT_BASE}/+/availability`, `{MQTT_BASE}/+/state`, `{MQTT_BASE}/+/json_attributes`, `{MQTT_BASE}/+/raw`, `{MQTT_BASE}/+/command_result`, and `{HA_DISCOVERY_PREFIX}/vacuum/+/config`. The availability write permission also covers the broker's last will. The `+` stands for one robot serial or discovery object ID. Home Assistant needs the corresponding subscribe permissions for discovery and state and publish permissions for commands.
|
||||||
|
|
||||||
|
The vacuum exposes start, stop, dock, spot, locate, and fan speed. It does not advertise pause, resume, mapping, segment cleaning, or a separate battery feature. State may be `idle`, `cleaning`, `returning`, `docked`, or `error`; attributes include `battery_level`, `side_brush`, `main_brush`, `filter`, `lifespan_total`, `clean_type`, `charge_state`, `last_error`, `last_command_error`, and `schedules`. Unknown readings start as JSON `null`. The bridge distinguishes a robot fault (`last_error`) from a failed or rejected command (`last_command_error`).
|
||||||
|
|
||||||
|
For a direct MQTT client, publish to `ecovacs/<serial>/command` with payload `start` (substitute your topic root and serial). Home Assistant normally publishes these commands through the discovered vacuum entity.
|
||||||
|
|
||||||
|
### Extension commands
|
||||||
|
|
||||||
|
The `send_command` topic accepts JSON objects with a `command` field; its maximum payload is **4096 bytes**. Examples below are MQTT payloads, not shell commands:
|
||||||
|
|
||||||
|
```text
|
||||||
|
{"command":"clean","clean_type":"border"}
|
||||||
|
{"command":"move","action":"SpinLeft"}
|
||||||
|
{"command":"cancel_return"}
|
||||||
|
{"command":"set_time"}
|
||||||
|
{"command":"get_sched"}
|
||||||
|
{"command":"add_sched","name":"weekday","on":"1","time":"21:30","repeat":"0111110","clean_type":"auto"}
|
||||||
|
{"command":"mod_sched","name":"weekday","on":"0","time":"21:30","repeat":"0111110","clean_type":"auto"}
|
||||||
|
{"command":"del_sched","name":"weekday"}
|
||||||
|
```
|
||||||
|
|
||||||
|
`clean_type` for `clean` may be `auto`, `border`, `spot`, or `singleRoom`. `move` accepts `forward`, `SpinLeft`, `SpinRight`, `TurnAround`, or `stop`; it does not automatically send a stop after a movement. The schedule `repeat` mask has seven `0`/`1` characters in **Sunday through Saturday** order. Schedule mutation supports `clean_type: "auto"` only. `on` accepts `"1"`, `"0"`, `"true"`, or `"false"`; `name` must be 1–64 bytes, and `time` must have `HH:MM` form. `get_sched` refreshes the schedule attribute. The bridge also retrieves battery, cleaning state, charge state, speed, schedules, and three consumable values on each successful session announcement.
|
||||||
|
|
||||||
|
Current limitation: the JSON parser accepts `{"command":"get_status"}` and `{"command":"get_lifespan"}`, but their actor paths do not build outbound XMPP queries. Do not rely on these two manual refresh commands in this version. Their values still update from the automatic session queries and incoming reports.
|
||||||
|
|
||||||
|
`RAW_COMMANDS=true` enables `{"command":"raw","xml":"<ctl td=\"...\" id=\"...\">...</ctl>"}`. It is disabled by default. The bridge accepts a `<ctl>` root, re-encodes its inner XML, takes `td`, and generates its own command ID; it does not forward the supplied root attributes verbatim. Use it only with a clear understanding of the N95 protocol. `resume` and `backward` are explicitly rejected.
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
| Symptom | Check |
|
||||||
|
| --- | --- |
|
||||||
|
| Container exits with `invalid configuration` | Read the joined validation errors in `docker compose logs n95bridge`. Check required addresses, strict boolean spelling, distinct ports, and PEM path/readability. |
|
||||||
|
| Health check passes but no robot appears | Verify the robot's DHCP DNS server and its `lbo.ecouser.net` answer, then test both lookup responses and the firmware 404 from the LAN. Reboot the robot after changing its cached endpoint. |
|
||||||
|
| `sasl authenticated` never appears | Check robot reachability to the published TCP 5223 port and whether its DNS and lookup requests reach this host. The XMPP listener uses plaintext. |
|
||||||
|
| Robot authenticates, but MQTT discovery is absent | Check for the `mqtt connect` log line, broker reachability from the container, broker credentials/TLS CA, and the Home Assistant discovery prefix. The MQTT client starts after robot READY. |
|
||||||
|
| Vacuum is present but `offline` | The robot session has not completed its announce ping, has disconnected, or has missed a keepalive response. The bridge pings about every 60 seconds and treats a missing response after 12 seconds as offline. |
|
||||||
|
| Values or commands do not update | Inspect `json_attributes.last_command_error`, then the nonretained `command_result` and `raw` topics if your broker permissions allow. Command payloads are exact and case sensitive. Offline commands are dropped rather than queued. |
|
||||||
|
| Time or schedules are wrong | Check the machine clock, `TZ`, and the Sunday-first schedule mask. |
|
||||||
|
|
||||||
|
The bridge logs to stderr. `LOG_LEVEL=debug` adds redacted XMPP stanza logging. The `raw` MQTT diagnostic topic can contain protocol XML, so grant access only to trusted MQTT clients. The bridge does not intentionally publish SASL authentication material there.
|
||||||
|
|
||||||
|
## Development and project documents
|
||||||
|
|
||||||
|
Run `go test ./...` from the repository root. The capture replay tests expect four `packetcapture-ix1.12-*.pcap` files at that root; these are local, untracked fixtures and may be absent from a fresh checkout. Without them, the capture tests fail even though the runtime build does not need them. A broker integration test is available with `go test -tags mqttintegration ./cmd/n95bridge` and `MQTT_TEST_URL` set to a test broker URL; it skips when that variable is unset. `MQTT_TEST_URL` is test-only, not a service setting.
|
||||||
|
|
||||||
|
Protocol and design details are in [the full N95 specification](docs/N95-FULL-SPECIFICATION.md), [packet capture analysis](docs/PCAP-ANALYSIS.md), [MQTT bridge mapping](docs/MQTT-BRIDGE.md), and [the milestone plan](docs/M1-MEGAPLAN.md). Those documents describe some intended or observed protocol behavior beyond the currently implemented runtime; the configuration and limitations above reflect the code in this repository.
|
||||||
|
|||||||
@@ -0,0 +1,109 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func findRepoFile(name string) (string, error) {
|
||||||
|
candidates := []string{
|
||||||
|
name,
|
||||||
|
filepath.Join("..", "..", name),
|
||||||
|
filepath.Join("..", name),
|
||||||
|
}
|
||||||
|
for _, p := range candidates {
|
||||||
|
if _, err := os.Stat(p); err == nil {
|
||||||
|
return p, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("file %q not found", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
func readRepoFile(t *testing.T, name string) string {
|
||||||
|
t.Helper()
|
||||||
|
path, err := findRepoFile(name)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("find file %s: %v", name, err)
|
||||||
|
}
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read file %s: %v", path, err)
|
||||||
|
}
|
||||||
|
return string(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComposeDNSAndAdvertiseIP(t *testing.T) {
|
||||||
|
content := readRepoFile(t, "docker-compose.yml")
|
||||||
|
|
||||||
|
// 1. Must publish 8005:8005, 8007:8007, 5223:5223.
|
||||||
|
requiredPorts := []string{"8005:8005", "8007:8007", "5223:5223"}
|
||||||
|
for _, port := range requiredPorts {
|
||||||
|
if !strings.Contains(content, port) {
|
||||||
|
t.Errorf("docker-compose.yml missing port mapping %q", port)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Must not use host networking.
|
||||||
|
if strings.Contains(content, "network_mode: host") || strings.Contains(content, "network_mode: \"host\"") {
|
||||||
|
t.Error("docker-compose.yml must not use network_mode: host")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. Comment block must mention lbo.ecouser.net and ADVERTISE_IP.
|
||||||
|
if !strings.Contains(content, "lbo.ecouser.net") {
|
||||||
|
t.Error("docker-compose.yml comment must contain lbo.ecouser.net")
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, "ADVERTISE_IP") {
|
||||||
|
t.Error("docker-compose.yml comment must contain ADVERTISE_IP")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 4. Must not contain an MQTT_PASSWORD value.
|
||||||
|
lines := strings.Split(content, "\n")
|
||||||
|
for _, line := range lines {
|
||||||
|
trimmed := strings.TrimSpace(line)
|
||||||
|
if strings.HasPrefix(trimmed, "#") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if strings.Contains(trimmed, "MQTT_PASSWORD") {
|
||||||
|
parts := strings.SplitN(trimmed, ":", 2)
|
||||||
|
if len(parts) == 2 && strings.TrimSpace(parts[1]) != "" {
|
||||||
|
t.Errorf("docker-compose.yml must not contain an MQTT_PASSWORD value, found: %q", trimmed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Log("Specification §15 item B1 verified and closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDockerfileNonRootStatic(t *testing.T) {
|
||||||
|
content := readRepoFile(t, "Dockerfile")
|
||||||
|
|
||||||
|
if !strings.Contains(content, "golang:1.27.1") {
|
||||||
|
t.Error("Dockerfile must specify golang:1.27.1")
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, "CGO_ENABLED=0") {
|
||||||
|
t.Error("Dockerfile must set CGO_ENABLED=0")
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, "gcr.io/distroless/static-debian12:nonroot") {
|
||||||
|
t.Error("Dockerfile must use gcr.io/distroless/static-debian12:nonroot")
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, "@sha256:") {
|
||||||
|
t.Error("Dockerfile must pin images by sha256 digest")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDockerignoreExcludesPcap(t *testing.T) {
|
||||||
|
content := readRepoFile(t, ".dockerignore")
|
||||||
|
|
||||||
|
if !strings.Contains(content, "*.pcap") {
|
||||||
|
t.Error(".dockerignore must match *.pcap")
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, "docs/") {
|
||||||
|
t.Error(".dockerignore must exclude docs/")
|
||||||
|
}
|
||||||
|
if !strings.Contains(content, ".env") {
|
||||||
|
t.Error(".dockerignore must exclude .env")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"os/signal"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ha"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/httpx"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/mqttbridge"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/xmpp"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
// shutdownBridge is assigned by phase 04.
|
||||||
|
shutdownBridge func(context.Context) error
|
||||||
|
// shutdownXMPP is assigned by phase 02.
|
||||||
|
shutdownXMPP func(context.Context) error
|
||||||
|
// shutdownTimeout is the graceful shutdown budget.
|
||||||
|
shutdownTimeout = 5 * time.Second
|
||||||
|
// logOutput is the destination for slog TextHandler.
|
||||||
|
logOutput io.Writer = os.Stderr
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
if len(os.Args) == 2 && os.Args[1] == "healthcheck" {
|
||||||
|
os.Exit(healthcheck())
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
|
if err := run(ctx, os.Environ()); err != nil {
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func run(ctx context.Context, environ []string) error {
|
||||||
|
cfg, err := config.Load(environ)
|
||||||
|
if err != nil {
|
||||||
|
slog.Error("invalid configuration", "err", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
slog.SetDefault(slog.New(slog.NewTextHandler(logOutput, &slog.HandlerOptions{
|
||||||
|
Level: cfg.LogLevel,
|
||||||
|
})))
|
||||||
|
slog.Info("starting with config", "config", cfg)
|
||||||
|
|
||||||
|
groupCtx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
|
||||||
|
// Create late-bound sinks to break the circular dependency.
|
||||||
|
var bridge *mqttbridge.Bridge
|
||||||
|
diagSink := func(d session.Diagnostic) {
|
||||||
|
if bridge != nil {
|
||||||
|
bridge.DiagnosticSink()(d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
traceSink := func(serial string, tr ctl.Trace) {
|
||||||
|
if bridge != nil {
|
||||||
|
bridge.TraceSink()(serial, tr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fleet := robot.NewFleet(groupCtx, cfg.ControllerJID, nil, traceSink, diagSink, nil)
|
||||||
|
var errBridge error
|
||||||
|
bridge, errBridge = mqttbridge.New(groupCtx, cfg, fleet.Submit)
|
||||||
|
if errBridge != nil {
|
||||||
|
return errBridge
|
||||||
|
}
|
||||||
|
|
||||||
|
robot.ParseSchedules = robot.ScheduleParserHook
|
||||||
|
robot.ParseLifespan = robot.LifespanParserHook
|
||||||
|
bridge.SendCommand = ha.SendCommandHook(cfg)
|
||||||
|
|
||||||
|
fleet.SetRepublishFactory(bridge.Republisher)
|
||||||
|
fleet.SetSendCommandFunc(bridge.DispatchSendCommand)
|
||||||
|
bus.Register(fleet)
|
||||||
|
bus.Register(bridge)
|
||||||
|
if shutdownBridge == nil {
|
||||||
|
shutdownBridge = bridge.Shutdown
|
||||||
|
}
|
||||||
|
xmppServer := xmpp.NewServer(cfg, registry, bus, nil, diagSink)
|
||||||
|
if shutdownXMPP == nil {
|
||||||
|
shutdownXMPP = xmppServer.Shutdown
|
||||||
|
}
|
||||||
|
|
||||||
|
errCh := make(chan error, 4)
|
||||||
|
go func() { errCh <- httpx.ServeLookup(groupCtx, cfg) }()
|
||||||
|
go func() { errCh <- httpx.ServeFirmware(groupCtx, cfg) }()
|
||||||
|
go func() { errCh <- httpx.ServeHealth(groupCtx, cfg) }()
|
||||||
|
go func() { errCh <- xmppServer.Serve(groupCtx) }()
|
||||||
|
|
||||||
|
var runErr error
|
||||||
|
select {
|
||||||
|
case runErr = <-errCh:
|
||||||
|
cancel()
|
||||||
|
case <-ctx.Done():
|
||||||
|
}
|
||||||
|
|
||||||
|
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), shutdownTimeout)
|
||||||
|
defer shutdownCancel()
|
||||||
|
if shutdownBridge != nil {
|
||||||
|
_ = shutdownBridge(shutdownCtx)
|
||||||
|
}
|
||||||
|
if shutdownXMPP != nil {
|
||||||
|
_ = shutdownXMPP(shutdownCtx)
|
||||||
|
}
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
if err := <-errCh; err != nil && runErr == nil {
|
||||||
|
runErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if shutdownCtx.Err() != nil {
|
||||||
|
slog.Info("shutdown budget expired")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if runErr != nil {
|
||||||
|
slog.Error("listener failed", "err", runErr)
|
||||||
|
}
|
||||||
|
return runErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func healthcheck() int {
|
||||||
|
cfg, err := config.Load(os.Environ())
|
||||||
|
if err != nil {
|
||||||
|
slog.Error("invalid configuration", "err", err)
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
url := fmt.Sprintf("http://127.0.0.1:%d/healthz", cfg.HealthPort)
|
||||||
|
client := &http.Client{Timeout: 2 * time.Second}
|
||||||
|
resp, err := client.Get(url)
|
||||||
|
if err != nil {
|
||||||
|
slog.Error("healthcheck request failed", "err", err)
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
slog.Error("healthcheck body read failed", "err", err)
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK || string(body) != "ok\n" {
|
||||||
|
slog.Error("healthcheck failed", "status", resp.StatusCode, "body", string(body))
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/httpx"
|
||||||
|
)
|
||||||
|
|
||||||
|
func withCleanEnv(t *testing.T, pairs ...string) {
|
||||||
|
t.Helper()
|
||||||
|
old := os.Environ()
|
||||||
|
t.Cleanup(func() {
|
||||||
|
os.Clearenv()
|
||||||
|
for _, e := range old {
|
||||||
|
if i := len("="); i > 0 {
|
||||||
|
// Split at first '='.
|
||||||
|
k, v, _ := splitEnv(e)
|
||||||
|
os.Setenv(k, v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
os.Clearenv()
|
||||||
|
for _, p := range pairs {
|
||||||
|
k, v, _ := splitEnv(p)
|
||||||
|
os.Setenv(k, v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func splitEnv(e string) (string, string, bool) {
|
||||||
|
for i := 0; i < len(e); i++ {
|
||||||
|
if e[i] == '=' {
|
||||||
|
return e[:i], e[i+1:], true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return e, "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthcheckSubcommandLive(t *testing.T) {
|
||||||
|
withCleanEnv(t,
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"HEALTH_PORT=18080",
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg, err := config.Load(os.Environ())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(done)
|
||||||
|
_ = httpx.ServeHealth(ctx, cfg)
|
||||||
|
}()
|
||||||
|
defer func() {
|
||||||
|
cancel()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("server did not stop")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
if got := healthcheck(); got != 0 {
|
||||||
|
t.Fatalf("healthcheck returned %d, want 0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthcheckSubcommandClosedPort(t *testing.T) {
|
||||||
|
withCleanEnv(t,
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"HEALTH_PORT=18080",
|
||||||
|
)
|
||||||
|
|
||||||
|
if got := healthcheck(); got != 1 {
|
||||||
|
t.Fatalf("healthcheck returned %d, want 1", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,190 @@
|
|||||||
|
//go:build mqttintegration
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
mqtt "github.com/eclipse/paho.mqtt.golang"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/mqttbridge"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMQTTIntegration(t *testing.T) {
|
||||||
|
testURL := os.Getenv("MQTT_TEST_URL")
|
||||||
|
if testURL == "" {
|
||||||
|
t.Skip("MQTT_TEST_URL is not set; skipping integration test")
|
||||||
|
}
|
||||||
|
|
||||||
|
u, err := url.Parse(testURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parse MQTT_TEST_URL: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
brokerHost := u.Hostname()
|
||||||
|
brokerPortStr := u.Port()
|
||||||
|
if brokerPortStr == "" {
|
||||||
|
brokerPortStr = "1883"
|
||||||
|
}
|
||||||
|
if _, err := strconv.Atoi(brokerPortStr); err != nil {
|
||||||
|
t.Fatalf("invalid broker port: %v", err)
|
||||||
|
}
|
||||||
|
brokerAddr := net.JoinHostPort(brokerHost, brokerPortStr)
|
||||||
|
|
||||||
|
// Start a local TCP proxy to simulate an abrupt network drop (unclean disconnect).
|
||||||
|
proxyLn, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("start proxy listener: %v", err)
|
||||||
|
}
|
||||||
|
defer proxyLn.Close()
|
||||||
|
proxyPort := proxyLn.Addr().(*net.TCPAddr).Port
|
||||||
|
|
||||||
|
var proxyMu sync.Mutex
|
||||||
|
var proxyConns []net.Conn
|
||||||
|
proxyDone := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer close(proxyDone)
|
||||||
|
for {
|
||||||
|
clientConn, err := proxyLn.Accept()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
brokerConn, err := net.Dial("tcp", brokerAddr)
|
||||||
|
if err != nil {
|
||||||
|
clientConn.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
proxyMu.Lock()
|
||||||
|
proxyConns = append(proxyConns, clientConn, brokerConn)
|
||||||
|
proxyMu.Unlock()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
_, _ = io.Copy(brokerConn, clientConn)
|
||||||
|
_ = brokerConn.Close()
|
||||||
|
}()
|
||||||
|
go func() {
|
||||||
|
_, _ = io.Copy(clientConn, brokerConn)
|
||||||
|
_ = clientConn.Close()
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
serial := "testserial123"
|
||||||
|
discoveryTopic := "homeassistant/vacuum/ecovacs_" + serial + "/config"
|
||||||
|
availabilityTopic := "ecovacs/" + serial + "/availability"
|
||||||
|
|
||||||
|
var (
|
||||||
|
mu sync.Mutex
|
||||||
|
discoveryMsg string
|
||||||
|
availMsg string
|
||||||
|
availHistory []string
|
||||||
|
)
|
||||||
|
discoveryCh := make(chan string, 10)
|
||||||
|
availCh := make(chan string, 10)
|
||||||
|
|
||||||
|
// Connect test subscriber directly to broker.
|
||||||
|
subOpts := mqtt.NewClientOptions()
|
||||||
|
subOpts.AddBroker(testURL)
|
||||||
|
subOpts.SetClientID("integration-test-subscriber")
|
||||||
|
subOpts.SetCleanSession(true)
|
||||||
|
subClient := mqtt.NewClient(subOpts)
|
||||||
|
if token := subClient.Connect(); !token.WaitTimeout(5*time.Second) || token.Error() != nil {
|
||||||
|
t.Fatalf("subscriber connect failed: %v", token.Error())
|
||||||
|
}
|
||||||
|
defer subClient.Disconnect(250)
|
||||||
|
|
||||||
|
subClient.Subscribe(discoveryTopic, 0, func(_ mqtt.Client, m mqtt.Message) {
|
||||||
|
mu.Lock()
|
||||||
|
discoveryMsg = string(m.Payload())
|
||||||
|
mu.Unlock()
|
||||||
|
discoveryCh <- string(m.Payload())
|
||||||
|
})
|
||||||
|
|
||||||
|
subClient.Subscribe(availabilityTopic, 0, func(_ mqtt.Client, m mqtt.Message) {
|
||||||
|
mu.Lock()
|
||||||
|
availMsg = string(m.Payload())
|
||||||
|
availHistory = append(availHistory, string(m.Payload()))
|
||||||
|
mu.Unlock()
|
||||||
|
availCh <- string(m.Payload())
|
||||||
|
})
|
||||||
|
|
||||||
|
// Create bridge pointing at the proxy port.
|
||||||
|
cfg := config.Config{
|
||||||
|
MQTTHost: "127.0.0.1",
|
||||||
|
MQTTPort: proxyPort,
|
||||||
|
MQTTBase: "ecovacs",
|
||||||
|
HADiscoveryPrefix: "homeassistant",
|
||||||
|
MQTTClientID: "testbridge",
|
||||||
|
ControllerJID: "n95bridge@ecouser.net/homeassistant",
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
submit := func(_ context.Context, _ string, _ robot.Command) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
bridge, err := mqttbridge.New(ctx, cfg, submit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("mqttbridge.New: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Trigger session announcement so bridge creates slot, connects, and publishes online.
|
||||||
|
bridge.AnnounceOK(session.ReadyEvent{
|
||||||
|
Serial: serial,
|
||||||
|
Generation: 1,
|
||||||
|
})
|
||||||
|
|
||||||
|
// 1. Assert retained discovery payload is published.
|
||||||
|
select {
|
||||||
|
case disc := <-discoveryCh:
|
||||||
|
if disc == "" {
|
||||||
|
t.Fatal("empty discovery payload received")
|
||||||
|
}
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for retained discovery payload")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Assert retained availability payload "online" is published.
|
||||||
|
select {
|
||||||
|
case avail := <-availCh:
|
||||||
|
if avail != "online" {
|
||||||
|
t.Fatalf("expected availability 'online', got: %q", avail)
|
||||||
|
}
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for retained availability online")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. Disconnect bridge without clean disconnect (kill proxy connections so broker detects drop).
|
||||||
|
proxyMu.Lock()
|
||||||
|
for _, c := range proxyConns {
|
||||||
|
_ = c.Close()
|
||||||
|
}
|
||||||
|
proxyMu.Unlock()
|
||||||
|
_ = proxyLn.Close()
|
||||||
|
|
||||||
|
// 4. Assert broker fires Last Will and publishes "offline".
|
||||||
|
select {
|
||||||
|
case avail := <-availCh:
|
||||||
|
if avail != "offline" {
|
||||||
|
t.Fatalf("expected LWT 'offline', got: %q", avail)
|
||||||
|
}
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for broker retained will offline")
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = discoveryMsg
|
||||||
|
_ = availMsg
|
||||||
|
}
|
||||||
@@ -0,0 +1,752 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/binary"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"regexp"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/httpx"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/xmpp"
|
||||||
|
)
|
||||||
|
|
||||||
|
const pcapMagicLittleEndian uint32 = 0xa1b2c3d4
|
||||||
|
|
||||||
|
var captureFilenames = []string{
|
||||||
|
"packetcapture-ix1.12-20260923211202.pcap",
|
||||||
|
"packetcapture-ix1.12-20260923214817.pcap",
|
||||||
|
"packetcapture-ix1.12-20260923223825.pcap",
|
||||||
|
"packetcapture-ix1.12-20260923230041.pcap",
|
||||||
|
}
|
||||||
|
|
||||||
|
type tcpSegment struct {
|
||||||
|
srcPort uint16
|
||||||
|
dstPort uint16
|
||||||
|
seq uint32
|
||||||
|
ack uint32
|
||||||
|
flags uint16
|
||||||
|
payload []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func readPCAPSegments(path string) ([]tcpSegment, error) {
|
||||||
|
f, err := os.Open(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer f.Close()
|
||||||
|
|
||||||
|
var gh struct {
|
||||||
|
Magic uint32
|
||||||
|
VersionMajor uint16
|
||||||
|
VersionMinor uint16
|
||||||
|
ThisZone int32
|
||||||
|
SigFigs uint32
|
||||||
|
SnapLen uint32
|
||||||
|
Network uint32
|
||||||
|
}
|
||||||
|
if err := binary.Read(f, binary.LittleEndian, &gh); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if gh.Magic != pcapMagicLittleEndian {
|
||||||
|
return nil, fmt.Errorf("unexpected magic: 0x%x", gh.Magic)
|
||||||
|
}
|
||||||
|
|
||||||
|
var segs []tcpSegment
|
||||||
|
for {
|
||||||
|
var ph struct {
|
||||||
|
TsSec uint32
|
||||||
|
TsUsec uint32
|
||||||
|
InclLen uint32
|
||||||
|
OrigLen uint32
|
||||||
|
}
|
||||||
|
err := binary.Read(f, binary.LittleEndian, &ph)
|
||||||
|
if err == io.EOF || errors.Is(err, io.ErrUnexpectedEOF) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
pkt := make([]byte, ph.InclLen)
|
||||||
|
if _, err := io.ReadFull(f, pkt); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ethernet header (14 bytes)
|
||||||
|
if len(pkt) < 14 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ethType := binary.BigEndian.Uint16(pkt[12:14])
|
||||||
|
if ethType != 0x0800 { // IPv4
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ipHdr := pkt[14:]
|
||||||
|
if len(ipHdr) < 20 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ihl := int(ipHdr[0]&0x0f) * 4
|
||||||
|
if len(ipHdr) < ihl {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
totLen := int(binary.BigEndian.Uint16(ipHdr[2:4]))
|
||||||
|
if totLen < ihl {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if totLen > len(ipHdr) {
|
||||||
|
totLen = len(ipHdr)
|
||||||
|
}
|
||||||
|
if ipHdr[9] != 6 { // TCP
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
tcpData := ipHdr[ihl:totLen]
|
||||||
|
if len(tcpData) < 20 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
srcPort := binary.BigEndian.Uint16(tcpData[0:2])
|
||||||
|
dstPort := binary.BigEndian.Uint16(tcpData[2:4])
|
||||||
|
seq := binary.BigEndian.Uint32(tcpData[4:8])
|
||||||
|
ack := binary.BigEndian.Uint32(tcpData[8:12])
|
||||||
|
offsetFlags := binary.BigEndian.Uint16(tcpData[12:14])
|
||||||
|
tcpLen := int((offsetFlags >> 12) & 0x0f) * 4
|
||||||
|
if len(tcpData) < tcpLen {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
flags := offsetFlags & 0x1ff
|
||||||
|
payload := tcpData[tcpLen:]
|
||||||
|
|
||||||
|
segs = append(segs, tcpSegment{
|
||||||
|
srcPort: srcPort,
|
||||||
|
dstPort: dstPort,
|
||||||
|
seq: seq,
|
||||||
|
ack: ack,
|
||||||
|
flags: flags,
|
||||||
|
payload: payload,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return segs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func reassembleClientFlows(segs []tcpSegment, dstPort uint16) map[uint16][]byte {
|
||||||
|
bySrc := make(map[uint16][]tcpSegment)
|
||||||
|
for _, s := range segs {
|
||||||
|
if s.dstPort == dstPort && len(s.payload) > 0 {
|
||||||
|
bySrc[s.srcPort] = append(bySrc[s.srcPort], s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
flows := make(map[uint16][]byte)
|
||||||
|
for sport, sl := range bySrc {
|
||||||
|
sort.Slice(sl, func(i, j int) bool {
|
||||||
|
return sl[i].seq < sl[j].seq
|
||||||
|
})
|
||||||
|
var data []byte
|
||||||
|
var nextSeq uint32
|
||||||
|
first := true
|
||||||
|
for _, s := range sl {
|
||||||
|
if first {
|
||||||
|
data = append(data, s.payload...)
|
||||||
|
nextSeq = s.seq + uint32(len(s.payload))
|
||||||
|
first = false
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if s.seq <= nextSeq {
|
||||||
|
overlap := int(nextSeq - s.seq)
|
||||||
|
if overlap < len(s.payload) {
|
||||||
|
data = append(data, s.payload[overlap:]...)
|
||||||
|
nextSeq += uint32(len(s.payload) - overlap)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
data = append(data, s.payload...)
|
||||||
|
nextSeq = s.seq + uint32(len(s.payload))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
flows[sport] = data
|
||||||
|
}
|
||||||
|
return flows
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCaptureFilesPresent(t *testing.T) {
|
||||||
|
for _, name := range captureFilenames {
|
||||||
|
p, err := findRepoFile(name)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("capture file %s not found: %v", name, err)
|
||||||
|
}
|
||||||
|
f, err := os.Open(p)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open %s: %v", p, err)
|
||||||
|
}
|
||||||
|
defer f.Close()
|
||||||
|
|
||||||
|
var magic uint32
|
||||||
|
if err := binary.Read(f, binary.LittleEndian, &magic); err != nil {
|
||||||
|
t.Fatalf("read magic from %s: %v", p, err)
|
||||||
|
}
|
||||||
|
if magic != pcapMagicLittleEndian {
|
||||||
|
t.Fatalf("file %s magic = 0x%x, want 0x%x", p, magic, pcapMagicLittleEndian)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCapture1LookupPair(t *testing.T) {
|
||||||
|
pcapPath, err := findRepoFile("packetcapture-ix1.12-20260923211202.pcap")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("pcap file not found: %v", err)
|
||||||
|
}
|
||||||
|
segs, err := readPCAPSegments(pcapPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read pcap: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
flows := reassembleClientFlows(segs, 8007)
|
||||||
|
if len(flows) != 2 {
|
||||||
|
t.Fatalf("expected 2 client lookup flows to port 8007, got %d", len(flows))
|
||||||
|
}
|
||||||
|
|
||||||
|
pLookup := getFreePort(t)
|
||||||
|
testAdvIP := "192.0.2.77"
|
||||||
|
cfg := config.Config{
|
||||||
|
BindAddress: "127.0.0.1",
|
||||||
|
AdvertiseIP: net.ParseIP(testAdvIP),
|
||||||
|
PortLookup: pLookup,
|
||||||
|
PortXMPP: 5223,
|
||||||
|
PortFirmware: 8005,
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
go func() {
|
||||||
|
_ = httpx.ServeLookup(ctx, cfg)
|
||||||
|
}()
|
||||||
|
time.Sleep(30 * time.Millisecond)
|
||||||
|
|
||||||
|
for sport, reqBytes := range flows {
|
||||||
|
conn, err := net.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", pLookup))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial lookup port %d: %v", pLookup, err)
|
||||||
|
}
|
||||||
|
_, err = conn.Write(reqBytes)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
t.Fatalf("write lookup request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
t.Fatalf("read lookup response: %v", err)
|
||||||
|
}
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
conn.Close()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read body: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("sport %d: status = %d, want 200", sport, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyStr := string(body)
|
||||||
|
if strings.Contains(bodyStr, " ") || strings.Contains(bodyStr, "\n") {
|
||||||
|
t.Fatalf("sport %d: response body is not compact JSON: %q", sport, bodyStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
var parsed struct {
|
||||||
|
Result string `json:"result"`
|
||||||
|
IP string `json:"ip"`
|
||||||
|
Port int `json:"port"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &parsed); err != nil {
|
||||||
|
t.Fatalf("sport %d: invalid JSON %q: %v", sport, bodyStr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsed.Result != "ok" {
|
||||||
|
t.Errorf("sport %d: result = %q, want ok", sport, parsed.Result)
|
||||||
|
}
|
||||||
|
if parsed.IP != testAdvIP {
|
||||||
|
t.Errorf("sport %d: ip = %q, want %s", sport, parsed.IP, testAdvIP)
|
||||||
|
}
|
||||||
|
|
||||||
|
if bytes.Contains(reqBytes, []byte("EcoMsgNew")) {
|
||||||
|
if parsed.Port != 5223 {
|
||||||
|
t.Errorf("EcoMsgNew port = %d, want 5223", parsed.Port)
|
||||||
|
}
|
||||||
|
} else if bytes.Contains(reqBytes, []byte("EcoUpdate")) {
|
||||||
|
if parsed.Port != 8005 {
|
||||||
|
t.Errorf("EcoUpdate port = %d, want 8005", parsed.Port)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
t.Fatalf("sport %d: request did not match EcoMsgNew or EcoUpdate: %s", sport, reqBytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCapture1Firmware404(t *testing.T) {
|
||||||
|
pcapPath, err := findRepoFile("packetcapture-ix1.12-20260923211202.pcap")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("pcap file not found: %v", err)
|
||||||
|
}
|
||||||
|
segs, err := readPCAPSegments(pcapPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read pcap: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
flows := reassembleClientFlows(segs, 8005)
|
||||||
|
if len(flows) != 1 {
|
||||||
|
t.Fatalf("expected 1 client firmware flow to port 8005, got %d", len(flows))
|
||||||
|
}
|
||||||
|
|
||||||
|
var reqBytes []byte
|
||||||
|
for _, b := range flows {
|
||||||
|
reqBytes = b
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
pFirmware := getFreePort(t)
|
||||||
|
cfg := config.Config{
|
||||||
|
BindAddress: "127.0.0.1",
|
||||||
|
PortFirmware: pFirmware,
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
go func() {
|
||||||
|
_ = httpx.ServeFirmware(ctx, cfg)
|
||||||
|
}()
|
||||||
|
time.Sleep(30 * time.Millisecond)
|
||||||
|
|
||||||
|
conn, err := net.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", pFirmware))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial firmware port %d: %v", pFirmware, err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
if _, err := conn.Write(reqBytes); err != nil {
|
||||||
|
t.Fatalf("write firmware request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read response: %v", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusNotFound {
|
||||||
|
t.Fatalf("status = %d, want 404", resp.StatusCode)
|
||||||
|
}
|
||||||
|
ct := resp.Header.Get("Content-Type")
|
||||||
|
if ct != "text/plain; charset=utf-8" {
|
||||||
|
t.Errorf("content-type = %q, want text/plain; charset=utf-8", ct)
|
||||||
|
}
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read body: %v", err)
|
||||||
|
}
|
||||||
|
if string(body) != "Not Found" {
|
||||||
|
t.Fatalf("body = %q, want Not Found", string(body))
|
||||||
|
}
|
||||||
|
if len(body) != 9 {
|
||||||
|
t.Fatalf("body len = %d, want 9", len(body))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type pcapXMPPClient struct {
|
||||||
|
t *testing.T
|
||||||
|
conn net.Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pcapXMPPClient) write(b []byte) {
|
||||||
|
if _, err := c.conn.Write(b); err != nil {
|
||||||
|
c.t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pcapXMPPClient) readUntil(delim string) string {
|
||||||
|
buf := make([]byte, 4096)
|
||||||
|
var accumulated string
|
||||||
|
for {
|
||||||
|
c.conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||||
|
n, err := c.conn.Read(buf)
|
||||||
|
if n > 0 {
|
||||||
|
accumulated += string(buf[:n])
|
||||||
|
if strings.Contains(accumulated, delim) {
|
||||||
|
return accumulated
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
c.t.Fatalf("readUntil %q failed after reading %q: %v", delim, accumulated, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCapture1Handshake(t *testing.T) {
|
||||||
|
pcapPath, err := findRepoFile("packetcapture-ix1.12-20260923211202.pcap")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("pcap file not found: %v", err)
|
||||||
|
}
|
||||||
|
segs, err := readPCAPSegments(pcapPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read pcap: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
flows := reassembleClientFlows(segs, 5223)
|
||||||
|
flow, ok := flows[10561]
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("client flow 10561 -> 5223 not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract SASL character data from auth element in the captured stream.
|
||||||
|
reAuth := regexp.MustCompile(`<auth[^>]*>([^<]+)</auth>`)
|
||||||
|
authMatch := reAuth.FindSubmatch(flow)
|
||||||
|
if authMatch == nil {
|
||||||
|
t.Fatal("auth element not found in captured stream")
|
||||||
|
}
|
||||||
|
saslChars := string(authMatch[1])
|
||||||
|
|
||||||
|
// Extract handshake stanzas from captured bytes.
|
||||||
|
streamOpen1 := []byte(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
authStanza := reAuth.Find(flow)
|
||||||
|
streamOpen2 := []byte(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
reBind := regexp.MustCompile(`<iq\s+type='set'\s+id='0'><bind[^>]*><resource>atom</resource></bind></iq>`)
|
||||||
|
bindStanza := reBind.Find(flow)
|
||||||
|
if bindStanza == nil {
|
||||||
|
t.Fatal("bind stanza not found in flow")
|
||||||
|
}
|
||||||
|
reSession := regexp.MustCompile(`<iq\s+type='set'\s+id='1'><session[^>]*/>\s*</iq>`)
|
||||||
|
sessionStanza := reSession.Find(flow)
|
||||||
|
if sessionStanza == nil {
|
||||||
|
t.Fatal("session stanza not found in flow")
|
||||||
|
}
|
||||||
|
rePresence := regexp.MustCompile(`<presence><status>hello world</status></presence>`)
|
||||||
|
presenceStanza := rePresence.Find(flow)
|
||||||
|
if presenceStanza == nil {
|
||||||
|
t.Fatal("presence stanza not found in flow")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Capture log records to ensure SASL character data is never logged.
|
||||||
|
logBuf := &safeBuffer{}
|
||||||
|
origLogOutput := logOutput
|
||||||
|
logOutput = logBuf
|
||||||
|
defer func() { logOutput = origLogOutput }()
|
||||||
|
|
||||||
|
pXMPP := getFreePort(t)
|
||||||
|
cfg := config.Config{
|
||||||
|
BindAddress: "127.0.0.1",
|
||||||
|
PortXMPP: pXMPP,
|
||||||
|
ControllerJID: "n95bridge@ecouser.net/homeassistant",
|
||||||
|
LogLevel: slog.LevelDebug,
|
||||||
|
}
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
server := xmpp.NewServer(cfg, registry, bus, nil, nil)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
go func() {
|
||||||
|
_ = server.Serve(ctx)
|
||||||
|
}()
|
||||||
|
time.Sleep(30 * time.Millisecond)
|
||||||
|
|
||||||
|
netConn, err := net.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", pXMPP))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial xmpp port: %v", err)
|
||||||
|
}
|
||||||
|
defer netConn.Close()
|
||||||
|
|
||||||
|
xc := &pcapXMPPClient{t: t, conn: netConn}
|
||||||
|
|
||||||
|
// 1. Initial stream open -> server features with SASL PLAIN
|
||||||
|
xc.write(streamOpen1)
|
||||||
|
feat1 := xc.readUntil("</stream:features>")
|
||||||
|
if !strings.Contains(feat1, "urn:ietf:params:xml:ns:xmpp-sasl") {
|
||||||
|
t.Fatalf("features 1 missing SASL namespace: %s", feat1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Auth from pcap -> server success
|
||||||
|
xc.write(authStanza)
|
||||||
|
succ := xc.readUntil("<success")
|
||||||
|
if !strings.Contains(succ, "<success") {
|
||||||
|
t.Fatalf("expected SASL success, got: %s", succ)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. Second stream open -> server features with bind
|
||||||
|
xc.write(streamOpen2)
|
||||||
|
feat2 := xc.readUntil("</stream:features>")
|
||||||
|
if !strings.Contains(feat2, "urn:ietf:params:xml:ns:xmpp-bind") {
|
||||||
|
t.Fatalf("features 2 missing bind: %s", feat2)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 4. Bind -> server bind result with JID shape {authcid}@{domain}/{resource}
|
||||||
|
xc.write(bindStanza)
|
||||||
|
bindRes := xc.readUntil("</iq>")
|
||||||
|
reJID := regexp.MustCompile(`<jid>([^<]+)</jid>`)
|
||||||
|
jidMatch := reJID.FindStringSubmatch(bindRes)
|
||||||
|
if len(jidMatch) < 2 {
|
||||||
|
t.Fatalf("bind result missing jid: %s", bindRes)
|
||||||
|
}
|
||||||
|
boundJID := jidMatch[1]
|
||||||
|
reExpectedJID := regexp.MustCompile(`^[A-Za-z0-9]+@[A-Za-z0-9.]+/atom$`)
|
||||||
|
if !reExpectedJID.MatchString(boundJID) {
|
||||||
|
t.Fatalf("bound JID shape mismatch: %q, want {authcid}@{domain}/{resource}", boundJID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 5. Session -> server result id="1"
|
||||||
|
xc.write(sessionStanza)
|
||||||
|
sessRes := xc.readUntil("/>")
|
||||||
|
if !strings.Contains(sessRes, `type="result"`) || !strings.Contains(sessRes, `id="1"`) {
|
||||||
|
t.Fatalf("session result mismatch: %s", sessRes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 6. Dummy presence with surrounding spaces: "> dummy </presence>"
|
||||||
|
xc.write(presenceStanza)
|
||||||
|
presRes := xc.readUntil("</presence>")
|
||||||
|
if !strings.Contains(presRes, "> dummy </presence>") {
|
||||||
|
t.Fatalf("dummy presence missing surrounding spaces, got: %s", presRes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Assert that SASL character data was NEVER logged.
|
||||||
|
logs := logBuf.String()
|
||||||
|
if strings.Contains(logs, saslChars) {
|
||||||
|
t.Fatal("SASL character data from capture was found in logs")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Also check that decoded SASL password does not appear in logs.
|
||||||
|
decoded, decErr := base64.StdEncoding.DecodeString(saslChars)
|
||||||
|
if decErr == nil {
|
||||||
|
parts := bytes.Split(decoded, []byte{0})
|
||||||
|
for _, part := range parts {
|
||||||
|
if len(part) > 0 && !strings.Contains(boundJID, string(part)) {
|
||||||
|
if strings.Contains(logs, string(part)) {
|
||||||
|
t.Fatal("decoded SASL password was found in logs")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCapture1BareBattery(t *testing.T) {
|
||||||
|
pcapPath, err := findRepoFile("packetcapture-ix1.12-20260923211202.pcap")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("pcap file not found: %v", err)
|
||||||
|
}
|
||||||
|
segs, err := readPCAPSegments(pcapPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read pcap: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
flows := reassembleClientFlows(segs, 5223)
|
||||||
|
flow, ok := flows[10561]
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("flow 10561 not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
reBare := regexp.MustCompile(`(?s)<iq\s+to='[^']+'\s+type='set'\s+id='[^']+'>\s*<query\s+xmlns='com:ctl'>\s*<battery\s+power='(\d+)'\s*/>\s*</query>\s*</iq>`)
|
||||||
|
match := reBare.Find(flow)
|
||||||
|
if match == nil {
|
||||||
|
t.Fatal("bare battery stanza not found in capture 1")
|
||||||
|
}
|
||||||
|
|
||||||
|
in, err := ctl.Parse(match)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ctl.Parse failed: %v", err)
|
||||||
|
}
|
||||||
|
if in.Kind != ctl.KindBattery {
|
||||||
|
t.Fatalf("parsed kind = %v, want KindBattery", in.Kind)
|
||||||
|
}
|
||||||
|
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
if err := robot.Apply(&snap, in, ""); err != nil {
|
||||||
|
t.Fatalf("Apply bare battery failed: %v", err)
|
||||||
|
}
|
||||||
|
if snap.Attributes.BatteryLevel == nil || *snap.Attributes.BatteryLevel != 76 {
|
||||||
|
t.Fatalf("battery_level = %v, want 76", snap.Attributes.BatteryLevel)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify actor writes no IQ result for bare battery and updates battery_level.
|
||||||
|
var mu sync.Mutex
|
||||||
|
var sentStanzas []string
|
||||||
|
var lastSnap robot.Snapshot
|
||||||
|
sendFn := func(b []byte) error {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
sentStanzas = append(sentStanzas, string(b))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
pubFn := func(_ context.Context, s robot.Snapshot) {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
lastSnap = s
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
testBotJID := "E2001229192001911354@155.ecorobot.net/atom"
|
||||||
|
actor := robot.NewActor(ctx, testBotJID, "E2001229192001911354", "controller@ecouser.net/ha", pubFn, nil, nil, nil)
|
||||||
|
actor.SessionReady(session.ReadyEvent{
|
||||||
|
Generation: 1,
|
||||||
|
JID: testBotJID,
|
||||||
|
Serial: "E2001229192001911354",
|
||||||
|
Send: sendFn,
|
||||||
|
})
|
||||||
|
actor.Stanza(session.StanzaEvent{
|
||||||
|
Generation: 1,
|
||||||
|
JID: testBotJID,
|
||||||
|
Serial: "E2001229192001911354",
|
||||||
|
Stanza: match,
|
||||||
|
})
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
if len(sentStanzas) != 0 {
|
||||||
|
t.Fatalf("bare battery produced %d outbound stanzas, want 0 (no iq result)", len(sentStanzas))
|
||||||
|
}
|
||||||
|
if lastSnap.Attributes.BatteryLevel == nil || *lastSnap.Attributes.BatteryLevel != 76 {
|
||||||
|
t.Fatalf("battery_level in actor snapshot = %v, want 76", lastSnap.Attributes.BatteryLevel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCapture2Sched2(t *testing.T) {
|
||||||
|
pcapPath, err := findRepoFile("packetcapture-ix1.12-20260923214817.pcap")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("pcap file not found: %v", err)
|
||||||
|
}
|
||||||
|
segs, err := readPCAPSegments(pcapPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read pcap: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
flows := reassembleClientFlows(segs, 5223)
|
||||||
|
flow, ok := flows[29867]
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("flow 29867 not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
reSched2 := regexp.MustCompile(`(?s)<iq\s+to='[^']+'\s+type='set'\s+id='[^']+'>\s*<query\s+xmlns='com:ctl'>\s*<ctl\s+td='Sched2'>\s*<s\s+n='[^']+'[^>]*>.*?</s>\s*</ctl>\s*</query>\s*</iq>`)
|
||||||
|
match := reSched2.Find(flow)
|
||||||
|
if match == nil {
|
||||||
|
t.Fatal("Sched2 stanza not found in capture 2")
|
||||||
|
}
|
||||||
|
|
||||||
|
in, err := ctl.Parse(match)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ctl.Parse failed: %v", err)
|
||||||
|
}
|
||||||
|
if in.TD != "Sched2" {
|
||||||
|
t.Fatalf("in.TD = %q, want Sched2", in.TD)
|
||||||
|
}
|
||||||
|
|
||||||
|
robot.ParseSchedules = robot.ScheduleParserHook
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
snap.Attributes.Schedules = []robot.Schedule{
|
||||||
|
{Name: "old_sched_1", Time: "08:00", On: false},
|
||||||
|
{Name: "old_sched_2", Time: "09:00", On: true},
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := robot.Apply(&snap, in, ""); err != nil {
|
||||||
|
t.Fatalf("Apply Sched2: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(snap.Attributes.Schedules) != 1 {
|
||||||
|
t.Fatalf("expected schedules slice replaced with 1 entry, got %d", len(snap.Attributes.Schedules))
|
||||||
|
}
|
||||||
|
s := snap.Attributes.Schedules[0]
|
||||||
|
if s.Name != "17901966021514" {
|
||||||
|
t.Errorf("schedule Name = %q, want 17901966021514", s.Name)
|
||||||
|
}
|
||||||
|
if s.Time != "19:30" {
|
||||||
|
t.Errorf("schedule Time = %q, want 19:30", s.Time)
|
||||||
|
}
|
||||||
|
if !s.On {
|
||||||
|
t.Errorf("schedule On = false, want true")
|
||||||
|
}
|
||||||
|
if s.Repeat != "0001000" {
|
||||||
|
t.Errorf("schedule Repeat = %q, want 0001000", s.Repeat)
|
||||||
|
}
|
||||||
|
if s.Flag != "p" {
|
||||||
|
t.Errorf("schedule Flag = %q, want p", s.Flag)
|
||||||
|
}
|
||||||
|
if s.Action.Type != "auto" {
|
||||||
|
t.Errorf("schedule Action.Type = %q, want auto", s.Action.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCapture4ReconnectIsNotASleep tests the reconnect sequence documented in
|
||||||
|
// PCAP-ANALYSIS.md §11.1 and N95-FULL-SPECIFICATION.md §7:
|
||||||
|
//
|
||||||
|
// Timeline from packetcapture-ix1.12-20260923230041.pcap:
|
||||||
|
// - The last healthy bot ping was answered.
|
||||||
|
// - 120 seconds later the next bot ping was black-holed. The same TCP segment was
|
||||||
|
// retransmitted at +0.668, +2.342, +5.368, +11.426, +23.468, +47.666, and +95.863 seconds.
|
||||||
|
// - An earlier incomplete series in packetcapture-ix1.12-20260923223825.pcap used
|
||||||
|
// +0.918, +2.999, +7.016, +15.043, +31.116, and +63.289 seconds.
|
||||||
|
// - The robot sent FIN at +120.002 seconds from the original ping, with no XMPP stream
|
||||||
|
// close and without waiting for FIN-ACK.
|
||||||
|
// - A new TCP connection opened 4.999 seconds after that FIN. SYN to dummy presence took
|
||||||
|
// 0.455 seconds. There was no DNS, lookup, firmware check, or XEP-0198 resume.
|
||||||
|
//
|
||||||
|
// The bridge's deadline is 12 seconds (xmpp.PingResultTimeout). Reconnect handling
|
||||||
|
// accepts the down-then-ready transition without sleeping or delay.
|
||||||
|
func TestCapture4ReconnectIsNotASleep(t *testing.T) {
|
||||||
|
if xmpp.PingResultTimeout != 12*time.Second {
|
||||||
|
t.Fatalf("xmpp.PingResultTimeout = %v, want 12s", xmpp.PingResultTimeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
|
||||||
|
testJID := "E2001229192001911354@155.ecorobot.net/atom"
|
||||||
|
serial := "E2001229192001911354"
|
||||||
|
|
||||||
|
// Register initial generation.
|
||||||
|
gen1, replaced := registry.Bind(testJID, nil, nil)
|
||||||
|
if gen1 != 1 || replaced {
|
||||||
|
t.Fatalf("initial bind: gen=%d replaced=%v", gen1, replaced)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Trigger SessionDown (simulating TCP close / dead path detection).
|
||||||
|
bus.SessionDown(session.DownEvent{
|
||||||
|
Serial: serial,
|
||||||
|
JID: testJID,
|
||||||
|
Generation: 1,
|
||||||
|
Reason: "tcp-fin",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Robot connects fresh socket immediately and completes re-bind to READY.
|
||||||
|
gen2, replaced := registry.Bind(testJID, nil, nil)
|
||||||
|
if gen2 != 2 || !replaced {
|
||||||
|
t.Fatalf("reconnect bind: gen=%d replaced=%v", gen2, replaced)
|
||||||
|
}
|
||||||
|
|
||||||
|
bus.SessionReady(session.ReadyEvent{
|
||||||
|
Serial: serial,
|
||||||
|
JID: testJID,
|
||||||
|
Generation: 2,
|
||||||
|
})
|
||||||
|
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
if elapsed > 500*time.Millisecond {
|
||||||
|
t.Fatalf("reconnect took %v, expected in well under a second (not a sleep)", elapsed)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,216 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func getFreePort(t *testing.T) int {
|
||||||
|
t.Helper()
|
||||||
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get free port: %v", err)
|
||||||
|
}
|
||||||
|
defer ln.Close()
|
||||||
|
return ln.Addr().(*net.TCPAddr).Port
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGracefulStopPublishesOffline(t *testing.T) {
|
||||||
|
pLookup := getFreePort(t)
|
||||||
|
pFirmware := getFreePort(t)
|
||||||
|
pXMPP := getFreePort(t)
|
||||||
|
pHealth := getFreePort(t)
|
||||||
|
|
||||||
|
env := []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
fmt.Sprintf("PORT_LOOKUP=%d", pLookup),
|
||||||
|
fmt.Sprintf("PORT_FIRMWARE=%d", pFirmware),
|
||||||
|
fmt.Sprintf("PORT_XMPP=%d", pXMPP),
|
||||||
|
fmt.Sprintf("HEALTH_PORT=%d", pHealth),
|
||||||
|
}
|
||||||
|
|
||||||
|
var mu sync.Mutex
|
||||||
|
var order []string
|
||||||
|
var bridgeDeadline time.Time
|
||||||
|
var bridgeHasDeadline bool
|
||||||
|
|
||||||
|
fakeShutdownBridge := func(ctx context.Context) error {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
dl, ok := ctx.Deadline()
|
||||||
|
bridgeHasDeadline = ok
|
||||||
|
bridgeDeadline = dl
|
||||||
|
order = append(order, "bridge_published_retained_offline")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fakeShutdownXMPP := func(ctx context.Context) error {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
order = append(order, "xmpp_stream_closed")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
shutdownBridge = fakeShutdownBridge
|
||||||
|
shutdownXMPP = fakeShutdownXMPP
|
||||||
|
t.Cleanup(func() {
|
||||||
|
shutdownBridge = nil
|
||||||
|
shutdownXMPP = nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
runDone := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
runDone <- run(ctx, env)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Wait for health endpoint to become ready.
|
||||||
|
healthURL := fmt.Sprintf("http://127.0.0.1:%d/healthz", pHealth)
|
||||||
|
ready := false
|
||||||
|
for i := 0; i < 50; i++ {
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
resp, err := http.Get(healthURL)
|
||||||
|
if err == nil {
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode == http.StatusOK {
|
||||||
|
ready = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !ready {
|
||||||
|
t.Fatal("servers did not become ready")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Trigger graceful stop (simulating SIGTERM via root context cancel).
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case err := <-runDone:
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("run returned error on graceful shutdown: %v", err)
|
||||||
|
}
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("run did not exit within 5 second shutdown budget")
|
||||||
|
}
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
if len(order) != 2 {
|
||||||
|
t.Fatalf("expected 2 shutdown actions, got %d: %v", len(order), order)
|
||||||
|
}
|
||||||
|
if order[0] != "bridge_published_retained_offline" {
|
||||||
|
t.Errorf("expected step 1 to be bridge publishing offline, got: %s", order[0])
|
||||||
|
}
|
||||||
|
if order[1] != "xmpp_stream_closed" {
|
||||||
|
t.Errorf("expected step 2 to be xmpp stream closed, got: %s", order[1])
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bridgeHasDeadline {
|
||||||
|
t.Error("shutdownBridge context did not have a timeout deadline")
|
||||||
|
} else {
|
||||||
|
remaining := time.Until(bridgeDeadline)
|
||||||
|
if remaining > 5*time.Second {
|
||||||
|
t.Errorf("shutdown timeout deadline > 5s: %v", remaining)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type safeBuffer struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
buf bytes.Buffer
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *safeBuffer) Write(p []byte) (n int, err error) {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
return b.buf.Write(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *safeBuffer) String() string {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
return b.buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShutdownBudgetExpiredLogsAndExitsZero(t *testing.T) {
|
||||||
|
origTimeout := shutdownTimeout
|
||||||
|
shutdownTimeout = 50 * time.Millisecond
|
||||||
|
defer func() { shutdownTimeout = origTimeout }()
|
||||||
|
|
||||||
|
logBuf := &safeBuffer{}
|
||||||
|
origLogOutput := logOutput
|
||||||
|
logOutput = logBuf
|
||||||
|
defer func() { logOutput = origLogOutput }()
|
||||||
|
|
||||||
|
pLookup := getFreePort(t)
|
||||||
|
pFirmware := getFreePort(t)
|
||||||
|
pXMPP := getFreePort(t)
|
||||||
|
pHealth := getFreePort(t)
|
||||||
|
|
||||||
|
env := []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
fmt.Sprintf("PORT_LOOKUP=%d", pLookup),
|
||||||
|
fmt.Sprintf("PORT_FIRMWARE=%d", pFirmware),
|
||||||
|
fmt.Sprintf("PORT_XMPP=%d", pXMPP),
|
||||||
|
fmt.Sprintf("HEALTH_PORT=%d", pHealth),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Slow bridge shutdown that exceeds the budget.
|
||||||
|
shutdownBridge = func(ctx context.Context) error {
|
||||||
|
<-ctx.Done() // wait for timeout
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
|
shutdownXMPP = func(ctx context.Context) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
shutdownBridge = nil
|
||||||
|
shutdownXMPP = nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
runDone := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
runDone <- run(ctx, env)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Wait for health endpoint.
|
||||||
|
healthURL := fmt.Sprintf("http://127.0.0.1:%d/healthz", pHealth)
|
||||||
|
for i := 0; i < 50; i++ {
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
resp, err := http.Get(healthURL)
|
||||||
|
if err == nil {
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode == http.StatusOK {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case err := <-runDone:
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected exit 0 (nil error) on budget expiration, got: %v", err)
|
||||||
|
}
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("run did not return within expected window")
|
||||||
|
}
|
||||||
|
|
||||||
|
out := logBuf.String()
|
||||||
|
if !strings.Contains(out, "shutdown budget expired") {
|
||||||
|
t.Fatalf("expected log 'shutdown budget expired', got: %s", out)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
# Operator runbook & network contract:
|
||||||
|
# 1. The robot's DNS resolver must answer lbo.ecouser.net with an A record for ADVERTISE_IP.
|
||||||
|
# 2. This container does not serve DNS.
|
||||||
|
# 3. ADVERTISE_IP is the Docker host's reserved LAN address, not BIND_ADDRESS and not a container bridge address.
|
||||||
|
# 4. That address has to remain reachable for the robot's powered-on lifetime (N95-FULL-SPECIFICATION.md §3.4).
|
||||||
|
# 5. Broker credentials (MQTT_USERNAME, MQTT_PASSWORD) should be supplied from an optional gitignored .env file.
|
||||||
|
# Do not embed credentials in this file.
|
||||||
|
|
||||||
|
services:
|
||||||
|
n95bridge:
|
||||||
|
build: .
|
||||||
|
restart: unless-stopped
|
||||||
|
ports:
|
||||||
|
- "8005:8005"
|
||||||
|
- "8007:8007"
|
||||||
|
- "5223:5223"
|
||||||
|
- "8080:8080"
|
||||||
|
environment:
|
||||||
|
PORT_FIRMWARE: "8005"
|
||||||
|
PORT_LOOKUP: "8007"
|
||||||
|
PORT_XMPP: "5223"
|
||||||
|
BIND_ADDRESS: "0.0.0.0"
|
||||||
|
ADVERTISE_IP: "192.0.2.10"
|
||||||
|
MQTT_HOST: "mqtt.example.invalid"
|
||||||
|
MQTT_TLS: "false"
|
||||||
|
MQTT_BASE: "ecovacs"
|
||||||
|
HA_DISCOVERY_PREFIX: "homeassistant"
|
||||||
|
CONTROLLER_JID: "n95bridge@ecouser.net/homeassistant"
|
||||||
|
RAW_COMMANDS: "false"
|
||||||
|
LOG_LEVEL: "info"
|
||||||
|
HEALTH_PORT: "8080"
|
||||||
|
TZ: "Europe/London"
|
||||||
|
# Optional uncommitted credentials file:
|
||||||
|
# env_file:
|
||||||
|
# - .env
|
||||||
|
healthcheck:
|
||||||
|
test: ["CMD", "/n95bridge", "healthcheck"]
|
||||||
|
interval: 10s
|
||||||
|
timeout: 3s
|
||||||
|
retries: 3
|
||||||
|
start_period: 5s
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
module git.i3omb.com/gronod/ha-n95-local-control
|
||||||
|
|
||||||
|
go 1.27.1
|
||||||
|
|
||||||
|
require github.com/eclipse/paho.mqtt.golang v1.5.1
|
||||||
|
|
||||||
|
require (
|
||||||
|
github.com/gorilla/websocket v1.5.3 // indirect
|
||||||
|
golang.org/x/net v0.44.0 // indirect
|
||||||
|
golang.org/x/sync v0.17.0 // indirect
|
||||||
|
)
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
github.com/eclipse/paho.mqtt.golang v1.5.1 h1:/VSOv3oDLlpqR2Epjn1Q7b2bSTplJIeV2ISgCl2W7nE=
|
||||||
|
github.com/eclipse/paho.mqtt.golang v1.5.1/go.mod h1:1/yJCneuyOoCOzKSsOTUc0AJfpsItBGWvYpBLimhArU=
|
||||||
|
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||||
|
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||||
|
golang.org/x/net v0.44.0 h1:evd8IRDyfNBMBTTY5XRF1vaZlD+EmWx6x8PkhR04H/I=
|
||||||
|
golang.org/x/net v0.44.0/go.mod h1:ECOoLqd5U3Lhyeyo/QDCEVQ4sNgYsqvCZ722XogGieY=
|
||||||
|
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
|
||||||
|
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||||
@@ -0,0 +1,375 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/pem"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"regexp"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultPortFirmware = 8005
|
||||||
|
defaultPortLookup = 8007
|
||||||
|
defaultPortXMPP = 5223
|
||||||
|
defaultHealthPort = 8080
|
||||||
|
defaultBindAddress = "0.0.0.0"
|
||||||
|
defaultMQTTBase = "ecovacs"
|
||||||
|
defaultHADiscovery = "homeassistant"
|
||||||
|
defaultControllerJID = "n95bridge@ecouser.net/homeassistant"
|
||||||
|
defaultLogLevel = "info"
|
||||||
|
defaultMQTTTLSPort = 8883
|
||||||
|
defaultMQTTPlainPort = 1883
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
clientIDRe = regexp.MustCompile(`^[A-Za-z0-9_-]{1,64}$`)
|
||||||
|
topicPartRe = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// Config holds every variable from the frozen environment table.
|
||||||
|
// MQTTPassword is never emitted by String, Format, GoString, or LogValue.
|
||||||
|
type Config struct {
|
||||||
|
PortFirmware int
|
||||||
|
PortLookup int
|
||||||
|
PortXMPP int
|
||||||
|
HealthPort int
|
||||||
|
BindAddress string
|
||||||
|
AdvertiseIP net.IP
|
||||||
|
MQTTHost string
|
||||||
|
MQTTPort int
|
||||||
|
MQTTTLS bool
|
||||||
|
MQTTCAFile string
|
||||||
|
MQTTUsername string
|
||||||
|
MQTTPassword string
|
||||||
|
MQTTClientID string
|
||||||
|
MQTTBase string
|
||||||
|
HADiscoveryPrefix string
|
||||||
|
ControllerJID string
|
||||||
|
RawCommands bool
|
||||||
|
LogLevel slog.Level
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load parses and validates environ, which must be in KEY=value form.
|
||||||
|
// Absent keys use defaults; empty values are invalid except for MQTT_PORT and
|
||||||
|
// MQTT_CLIENT_ID, which fall back to their defaults.
|
||||||
|
func Load(environ []string) (Config, error) {
|
||||||
|
env := make(map[string]string, len(environ))
|
||||||
|
present := make(map[string]bool, len(environ))
|
||||||
|
for _, e := range environ {
|
||||||
|
if i := strings.IndexByte(e, '='); i >= 0 {
|
||||||
|
key := e[:i]
|
||||||
|
env[key] = e[i+1:]
|
||||||
|
present[key] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var cfg Config
|
||||||
|
var errs []string
|
||||||
|
|
||||||
|
parsePort := func(key, raw string, def int, allowEmpty bool) int {
|
||||||
|
if raw == "" {
|
||||||
|
if allowEmpty {
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
errs = append(errs, fmt.Sprintf("%s must be an integer from 1 to 65535", key))
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
n, err := strconv.Atoi(raw)
|
||||||
|
if err != nil || n < 1 || n > 65535 {
|
||||||
|
errs = append(errs, fmt.Sprintf("%s must be an integer from 1 to 65535", key))
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// MQTT_TLS is parsed before applying the MQTT_PORT default.
|
||||||
|
mqttTLS := false
|
||||||
|
if !present["MQTT_TLS"] {
|
||||||
|
mqttTLS = false
|
||||||
|
} else if v := env["MQTT_TLS"]; v == "" {
|
||||||
|
errs = append(errs, "MQTT_TLS must be true or false")
|
||||||
|
} else {
|
||||||
|
b, ok := parseBoolStrict(v)
|
||||||
|
if !ok {
|
||||||
|
errs = append(errs, "MQTT_TLS must be true or false")
|
||||||
|
} else {
|
||||||
|
mqttTLS = b
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cfg.MQTTTLS = mqttTLS
|
||||||
|
|
||||||
|
cfg.PortFirmware = parsePort("PORT_FIRMWARE", env["PORT_FIRMWARE"], defaultPortFirmware, !present["PORT_FIRMWARE"])
|
||||||
|
cfg.PortLookup = parsePort("PORT_LOOKUP", env["PORT_LOOKUP"], defaultPortLookup, !present["PORT_LOOKUP"])
|
||||||
|
cfg.PortXMPP = parsePort("PORT_XMPP", env["PORT_XMPP"], defaultPortXMPP, !present["PORT_XMPP"])
|
||||||
|
cfg.HealthPort = parsePort("HEALTH_PORT", env["HEALTH_PORT"], defaultHealthPort, !present["HEALTH_PORT"])
|
||||||
|
|
||||||
|
if !present["BIND_ADDRESS"] {
|
||||||
|
cfg.BindAddress = defaultBindAddress
|
||||||
|
} else if env["BIND_ADDRESS"] == "" || net.ParseIP(env["BIND_ADDRESS"]) == nil {
|
||||||
|
errs = append(errs, "BIND_ADDRESS must be an IP address")
|
||||||
|
} else {
|
||||||
|
cfg.BindAddress = env["BIND_ADDRESS"]
|
||||||
|
}
|
||||||
|
|
||||||
|
advRaw := strings.TrimSpace(env["ADVERTISE_IP"])
|
||||||
|
if advRaw == "" {
|
||||||
|
errs = append(errs, "ADVERTISE_IP is required")
|
||||||
|
} else {
|
||||||
|
ip := net.ParseIP(advRaw)
|
||||||
|
if ip == nil || ip.To4() == nil || ip.To4().String() != advRaw {
|
||||||
|
errs = append(errs, "ADVERTISE_IP must be a literal IPv4 address")
|
||||||
|
} else if ip.Equal(net.IPv4zero) {
|
||||||
|
errs = append(errs, "ADVERTISE_IP must not be 0.0.0.0")
|
||||||
|
} else {
|
||||||
|
cfg.AdvertiseIP = ip.To4()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
hostRaw := env["MQTT_HOST"]
|
||||||
|
if hostRaw == "" {
|
||||||
|
errs = append(errs, "MQTT_HOST is required")
|
||||||
|
} else if strings.Contains(hostRaw, "://") {
|
||||||
|
errs = append(errs, "MQTT_HOST must be a host or address without a scheme or port")
|
||||||
|
} else if _, _, err := net.SplitHostPort(hostRaw); err == nil {
|
||||||
|
errs = append(errs, "MQTT_HOST must be a host or address without a scheme or port")
|
||||||
|
} else {
|
||||||
|
cfg.MQTTHost = hostRaw
|
||||||
|
}
|
||||||
|
|
||||||
|
mqttPortDefault := defaultMQTTPlainPort
|
||||||
|
if mqttTLS {
|
||||||
|
mqttPortDefault = defaultMQTTTLSPort
|
||||||
|
}
|
||||||
|
cfg.MQTTPort = parsePort("MQTT_PORT", env["MQTT_PORT"], mqttPortDefault, true)
|
||||||
|
|
||||||
|
cfg.MQTTCAFile = env["MQTT_CA_FILE"]
|
||||||
|
if cfg.MQTTCAFile != "" {
|
||||||
|
if err := validatePEMCertBundle(cfg.MQTTCAFile); err != nil {
|
||||||
|
errs = append(errs, "MQTT_CA_FILE must be a PEM certificate bundle")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg.MQTTUsername = env["MQTT_USERNAME"]
|
||||||
|
cfg.MQTTPassword = env["MQTT_PASSWORD"]
|
||||||
|
if cfg.MQTTPassword != "" && cfg.MQTTUsername == "" {
|
||||||
|
errs = append(errs, "MQTT_PASSWORD requires MQTT_USERNAME")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !present["MQTT_CLIENT_ID"] || env["MQTT_CLIENT_ID"] == "" {
|
||||||
|
cfg.MQTTClientID = defaultMQTTClientID()
|
||||||
|
} else if !clientIDRe.MatchString(env["MQTT_CLIENT_ID"]) {
|
||||||
|
errs = append(errs, "MQTT_CLIENT_ID must match ^[A-Za-z0-9_-]{1,64}$")
|
||||||
|
} else {
|
||||||
|
cfg.MQTTClientID = env["MQTT_CLIENT_ID"]
|
||||||
|
}
|
||||||
|
|
||||||
|
if !present["MQTT_BASE"] {
|
||||||
|
cfg.MQTTBase = defaultMQTTBase
|
||||||
|
} else if env["MQTT_BASE"] == "" || !topicPartRe.MatchString(env["MQTT_BASE"]) {
|
||||||
|
errs = append(errs, "MQTT_BASE must match ^[A-Za-z0-9_-]+$")
|
||||||
|
} else {
|
||||||
|
cfg.MQTTBase = env["MQTT_BASE"]
|
||||||
|
}
|
||||||
|
|
||||||
|
if !present["HA_DISCOVERY_PREFIX"] {
|
||||||
|
cfg.HADiscoveryPrefix = defaultHADiscovery
|
||||||
|
} else if env["HA_DISCOVERY_PREFIX"] == "" || !topicPartRe.MatchString(env["HA_DISCOVERY_PREFIX"]) {
|
||||||
|
errs = append(errs, "HA_DISCOVERY_PREFIX must match ^[A-Za-z0-9_-]+$")
|
||||||
|
} else {
|
||||||
|
cfg.HADiscoveryPrefix = env["HA_DISCOVERY_PREFIX"]
|
||||||
|
}
|
||||||
|
|
||||||
|
if !present["CONTROLLER_JID"] {
|
||||||
|
cfg.ControllerJID = defaultControllerJID
|
||||||
|
} else if env["CONTROLLER_JID"] == "" || !validControllerJID(env["CONTROLLER_JID"]) {
|
||||||
|
errs = append(errs, "CONTROLLER_JID must be local@domain/resource")
|
||||||
|
} else {
|
||||||
|
cfg.ControllerJID = env["CONTROLLER_JID"]
|
||||||
|
}
|
||||||
|
|
||||||
|
if !present["RAW_COMMANDS"] {
|
||||||
|
cfg.RawCommands = false
|
||||||
|
} else if env["RAW_COMMANDS"] == "" {
|
||||||
|
errs = append(errs, "RAW_COMMANDS must be true or false")
|
||||||
|
} else {
|
||||||
|
b, ok := parseBoolStrict(env["RAW_COMMANDS"])
|
||||||
|
if !ok {
|
||||||
|
errs = append(errs, "RAW_COMMANDS must be true or false")
|
||||||
|
} else {
|
||||||
|
cfg.RawCommands = b
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
levelRaw := env["LOG_LEVEL"]
|
||||||
|
if !present["LOG_LEVEL"] {
|
||||||
|
levelRaw = defaultLogLevel
|
||||||
|
}
|
||||||
|
switch strings.ToLower(levelRaw) {
|
||||||
|
case "debug":
|
||||||
|
cfg.LogLevel = slog.LevelDebug
|
||||||
|
case "info":
|
||||||
|
cfg.LogLevel = slog.LevelInfo
|
||||||
|
case "warn":
|
||||||
|
cfg.LogLevel = slog.LevelWarn
|
||||||
|
case "error":
|
||||||
|
cfg.LogLevel = slog.LevelError
|
||||||
|
default:
|
||||||
|
errs = append(errs, "LOG_LEVEL must be debug, info, warn, or error")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.PortFirmware != 0 && cfg.PortLookup != 0 && cfg.PortXMPP != 0 && cfg.HealthPort != 0 {
|
||||||
|
if cfg.PortFirmware == cfg.PortLookup || cfg.PortFirmware == cfg.PortXMPP || cfg.PortFirmware == cfg.HealthPort ||
|
||||||
|
cfg.PortLookup == cfg.PortXMPP || cfg.PortLookup == cfg.HealthPort || cfg.PortXMPP == cfg.HealthPort {
|
||||||
|
errs = append(errs, "listener ports must be distinct")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(errs) > 0 {
|
||||||
|
return Config{}, fmt.Errorf("%s", strings.Join(errs, "; "))
|
||||||
|
}
|
||||||
|
return cfg, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseBoolStrict(s string) (bool, bool) {
|
||||||
|
switch s {
|
||||||
|
case "true":
|
||||||
|
return true, true
|
||||||
|
case "false":
|
||||||
|
return false, true
|
||||||
|
}
|
||||||
|
return false, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func validatePEMCertBundle(path string) error {
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for {
|
||||||
|
block, rest := pem.Decode(data)
|
||||||
|
if block == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if block.Type == "CERTIFICATE" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
data = rest
|
||||||
|
}
|
||||||
|
return fmt.Errorf("no certificate block")
|
||||||
|
}
|
||||||
|
|
||||||
|
func validControllerJID(jid string) bool {
|
||||||
|
if strings.Count(jid, "@") != 1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
local, domainResource, _ := strings.Cut(jid, "@")
|
||||||
|
if local == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if strings.Count(domainResource, "/") != 1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
domain, resource, _ := strings.Cut(domainResource, "/")
|
||||||
|
return domain != "" && resource != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultMQTTClientID() string {
|
||||||
|
host, err := os.Hostname()
|
||||||
|
if err != nil {
|
||||||
|
host = "host"
|
||||||
|
}
|
||||||
|
if i := strings.IndexByte(host, '.'); i >= 0 {
|
||||||
|
host = host[:i]
|
||||||
|
}
|
||||||
|
host = strings.ToLower(host)
|
||||||
|
var b strings.Builder
|
||||||
|
for i := 0; i < len(host); i++ {
|
||||||
|
c := host[i]
|
||||||
|
if (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') {
|
||||||
|
b.WriteByte(c)
|
||||||
|
} else {
|
||||||
|
b.WriteByte('-')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
hostPart := b.String()
|
||||||
|
hostPart = collapseTrim(hostPart, '-')
|
||||||
|
if hostPart == "" {
|
||||||
|
hostPart = "host"
|
||||||
|
}
|
||||||
|
prefix := "n95bridge-"
|
||||||
|
result := prefix + hostPart
|
||||||
|
if len(result) > 64 {
|
||||||
|
maxHost := 64 - len(prefix)
|
||||||
|
if maxHost < 0 {
|
||||||
|
maxHost = 0
|
||||||
|
}
|
||||||
|
if len(hostPart) > maxHost {
|
||||||
|
hostPart = hostPart[len(hostPart)-maxHost:]
|
||||||
|
}
|
||||||
|
hostPart = strings.TrimLeft(hostPart, "-")
|
||||||
|
result = prefix + hostPart
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func collapseTrim(s string, ch byte) string {
|
||||||
|
var b strings.Builder
|
||||||
|
prev := byte(0)
|
||||||
|
for i := 0; i < len(s); i++ {
|
||||||
|
c := s[i]
|
||||||
|
if c == ch {
|
||||||
|
if prev == ch {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.WriteByte(c)
|
||||||
|
prev = c
|
||||||
|
}
|
||||||
|
return strings.Trim(b.String(), string(ch))
|
||||||
|
}
|
||||||
|
|
||||||
|
// String returns a redacted summary that omits MQTTPassword.
|
||||||
|
func (c Config) String() string {
|
||||||
|
var b strings.Builder
|
||||||
|
fmt.Fprintf(&b, "PortFirmware=%d PortLookup=%d PortXMPP=%d HealthPort=%d ", c.PortFirmware, c.PortLookup, c.PortXMPP, c.HealthPort)
|
||||||
|
fmt.Fprintf(&b, "BindAddress=%q AdvertiseIP=%s ", c.BindAddress, c.AdvertiseIP)
|
||||||
|
fmt.Fprintf(&b, "MQTTHost=%s MQTTPort=%d MQTTTLS=%t MQTTCAFile=%q MQTTUsername=%q ", c.MQTTHost, c.MQTTPort, c.MQTTTLS, c.MQTTCAFile, c.MQTTUsername)
|
||||||
|
fmt.Fprintf(&b, "MQTTClientID=%q MQTTBase=%q HADiscoveryPrefix=%q ControllerJID=%q ", c.MQTTClientID, c.MQTTBase, c.HADiscoveryPrefix, c.ControllerJID)
|
||||||
|
fmt.Fprintf(&b, "RawCommands=%t LogLevel=%s", c.RawCommands, c.LogLevel)
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Format redirects all fmt verbs to the redacted String.
|
||||||
|
func (c Config) Format(s fmt.State, verb rune) {
|
||||||
|
_, _ = fmt.Fprint(s, c.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// GoString returns the redacted String.
|
||||||
|
func (c Config) GoString() string { return c.String() }
|
||||||
|
|
||||||
|
// LogValue returns a slog group that omits MQTTPassword.
|
||||||
|
func (c Config) LogValue() slog.Value {
|
||||||
|
return slog.GroupValue(
|
||||||
|
slog.Int("PortFirmware", c.PortFirmware),
|
||||||
|
slog.Int("PortLookup", c.PortLookup),
|
||||||
|
slog.Int("PortXMPP", c.PortXMPP),
|
||||||
|
slog.Int("HealthPort", c.HealthPort),
|
||||||
|
slog.String("BindAddress", c.BindAddress),
|
||||||
|
slog.String("AdvertiseIP", c.AdvertiseIP.String()),
|
||||||
|
slog.String("MQTTHost", c.MQTTHost),
|
||||||
|
slog.Int("MQTTPort", c.MQTTPort),
|
||||||
|
slog.Bool("MQTTTLS", c.MQTTTLS),
|
||||||
|
slog.String("MQTTCAFile", c.MQTTCAFile),
|
||||||
|
slog.String("MQTTUsername", c.MQTTUsername),
|
||||||
|
slog.String("MQTTClientID", c.MQTTClientID),
|
||||||
|
slog.String("MQTTBase", c.MQTTBase),
|
||||||
|
slog.String("HADiscoveryPrefix", c.HADiscoveryPrefix),
|
||||||
|
slog.String("ControllerJID", c.ControllerJID),
|
||||||
|
slog.Bool("RawCommands", c.RawCommands),
|
||||||
|
slog.String("LogLevel", c.LogLevel.String()),
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,383 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/rsa"
|
||||||
|
"crypto/x509"
|
||||||
|
"crypto/x509/pkix"
|
||||||
|
"encoding/pem"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"math/big"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func validEnv() []string {
|
||||||
|
return []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadValidMinimal(t *testing.T) {
|
||||||
|
cfg, err := Load(validEnv())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.PortFirmware != defaultPortFirmware {
|
||||||
|
t.Errorf("PortFirmware = %d, want %d", cfg.PortFirmware, defaultPortFirmware)
|
||||||
|
}
|
||||||
|
if cfg.PortLookup != defaultPortLookup {
|
||||||
|
t.Errorf("PortLookup = %d, want %d", cfg.PortLookup, defaultPortLookup)
|
||||||
|
}
|
||||||
|
if cfg.PortXMPP != defaultPortXMPP {
|
||||||
|
t.Errorf("PortXMPP = %d, want %d", cfg.PortXMPP, defaultPortXMPP)
|
||||||
|
}
|
||||||
|
if cfg.HealthPort != defaultHealthPort {
|
||||||
|
t.Errorf("HealthPort = %d, want %d", cfg.HealthPort, defaultHealthPort)
|
||||||
|
}
|
||||||
|
if cfg.BindAddress != defaultBindAddress {
|
||||||
|
t.Errorf("BindAddress = %q, want %q", cfg.BindAddress, defaultBindAddress)
|
||||||
|
}
|
||||||
|
if cfg.MQTTPort != defaultMQTTPlainPort {
|
||||||
|
t.Errorf("MQTTPort = %d, want %d", cfg.MQTTPort, defaultMQTTPlainPort)
|
||||||
|
}
|
||||||
|
if cfg.MQTTTLS {
|
||||||
|
t.Error("MQTTTLS should be false")
|
||||||
|
}
|
||||||
|
if cfg.MQTTBase != defaultMQTTBase {
|
||||||
|
t.Errorf("MQTTBase = %q, want %q", cfg.MQTTBase, defaultMQTTBase)
|
||||||
|
}
|
||||||
|
if cfg.HADiscoveryPrefix != defaultHADiscovery {
|
||||||
|
t.Errorf("HADiscoveryPrefix = %q, want %q", cfg.HADiscoveryPrefix, defaultHADiscovery)
|
||||||
|
}
|
||||||
|
if cfg.ControllerJID != defaultControllerJID {
|
||||||
|
t.Errorf("ControllerJID = %q, want %q", cfg.ControllerJID, defaultControllerJID)
|
||||||
|
}
|
||||||
|
if cfg.RawCommands {
|
||||||
|
t.Error("RawCommands should be false")
|
||||||
|
}
|
||||||
|
if cfg.LogLevel != slog.LevelInfo {
|
||||||
|
t.Errorf("LogLevel = %v, want info", cfg.LogLevel)
|
||||||
|
}
|
||||||
|
if cfg.MQTTClientID == "" || !clientIDRe.MatchString(cfg.MQTTClientID) {
|
||||||
|
t.Errorf("MQTTClientID = %q, want hostname-derived id matching regex", cfg.MQTTClientID)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(cfg.MQTTClientID, "n95bridge-") {
|
||||||
|
t.Errorf("MQTTClientID = %q, want prefix n95bridge-", cfg.MQTTClientID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadPortsAndTLS(t *testing.T) {
|
||||||
|
cfg, err := Load([]string{
|
||||||
|
"PORT_FIRMWARE=9005",
|
||||||
|
"PORT_LOOKUP=9007",
|
||||||
|
"PORT_XMPP=9223",
|
||||||
|
"HEALTH_PORT=9080",
|
||||||
|
"BIND_ADDRESS=127.0.0.1",
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_TLS=true",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.PortFirmware != 9005 || cfg.PortLookup != 9007 || cfg.PortXMPP != 9223 || cfg.HealthPort != 9080 {
|
||||||
|
t.Errorf("ports mismatch: %+v", cfg)
|
||||||
|
}
|
||||||
|
if cfg.BindAddress != "127.0.0.1" {
|
||||||
|
t.Errorf("BindAddress = %q", cfg.BindAddress)
|
||||||
|
}
|
||||||
|
if cfg.MQTTPort != defaultMQTTTLSPort {
|
||||||
|
t.Errorf("MQTTPort with TLS = %d, want %d", cfg.MQTTPort, defaultMQTTTLSPort)
|
||||||
|
}
|
||||||
|
if !cfg.MQTTTLS {
|
||||||
|
t.Error("MQTTTLS should be true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadPortDefaultFollowsTLS(t *testing.T) {
|
||||||
|
cfg, err := Load([]string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_TLS=true",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.MQTTPort != 8883 {
|
||||||
|
t.Errorf("MQTTPort = %d, want 8883", cfg.MQTTPort)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err = Load([]string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_PORT=9999",
|
||||||
|
"MQTT_TLS=true",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.MQTTPort != 9999 {
|
||||||
|
t.Errorf("MQTTPort = %d, want 9999", cfg.MQTTPort)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadValidationMessages(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
badPEM := filepath.Join(tmp, "bad.pem")
|
||||||
|
if err := os.WriteFile(badPEM, []byte("not a certificate"), 0600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
goodPEM := filepath.Join(tmp, "good.pem")
|
||||||
|
writeTestCertificate(t, goodPEM)
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
env []string
|
||||||
|
wantAll []string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "ports invalid",
|
||||||
|
env: []string{
|
||||||
|
"PORT_FIRMWARE=abc",
|
||||||
|
"PORT_LOOKUP=0",
|
||||||
|
"PORT_XMPP=99999",
|
||||||
|
"HEALTH_PORT=-1",
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"PORT_FIRMWARE must be an integer from 1 to 65535",
|
||||||
|
"PORT_LOOKUP must be an integer from 1 to 65535",
|
||||||
|
"PORT_XMPP must be an integer from 1 to 65535",
|
||||||
|
"HEALTH_PORT must be an integer from 1 to 65535",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "advertise",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=0.0.0.0",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"ADVERTISE_IP must not be 0.0.0.0",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mqtt host",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=tcp://mqtt.example.invalid",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"MQTT_HOST must be a host or address without a scheme or port",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mqtt ca",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_CA_FILE=" + badPEM,
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"MQTT_CA_FILE must be a PEM certificate bundle",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "password without username",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_PASSWORD=secret",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"MQTT_PASSWORD requires MQTT_USERNAME",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "client id",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_CLIENT_ID=bad id!",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"MQTT_CLIENT_ID must match ^[A-Za-z0-9_-]{1,64}$",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "topics jid",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_BASE=",
|
||||||
|
"HA_DISCOVERY_PREFIX=home assistant",
|
||||||
|
"CONTROLLER_JID=nobody",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"MQTT_BASE must match ^[A-Za-z0-9_-]+$",
|
||||||
|
"HA_DISCOVERY_PREFIX must match ^[A-Za-z0-9_-]+$",
|
||||||
|
"CONTROLLER_JID must be local@domain/resource",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "log and raw",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"RAW_COMMANDS=maybe",
|
||||||
|
"LOG_LEVEL=verbose",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"RAW_COMMANDS must be true or false",
|
||||||
|
"LOG_LEVEL must be debug, info, warn, or error",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "duplicate ports",
|
||||||
|
env: []string{
|
||||||
|
"PORT_FIRMWARE=8005",
|
||||||
|
"PORT_LOOKUP=8005",
|
||||||
|
"PORT_XMPP=5223",
|
||||||
|
"HEALTH_PORT=8080",
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
},
|
||||||
|
wantAll: []string{
|
||||||
|
"listener ports must be distinct",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ca ok",
|
||||||
|
env: []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"MQTT_CA_FILE=" + goodPEM,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
_, err := Load(tc.env)
|
||||||
|
if len(tc.wantAll) == 0 {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected error containing %q, got nil", tc.wantAll)
|
||||||
|
}
|
||||||
|
got := err.Error()
|
||||||
|
for _, want := range tc.wantAll {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("error %q does not contain %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadAdvertiseIPValidation(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
value string
|
||||||
|
wantErr string
|
||||||
|
}{
|
||||||
|
{"missing", "", "ADVERTISE_IP is required"},
|
||||||
|
{"not v4", "::1", "ADVERTISE_IP must be a literal IPv4 address"},
|
||||||
|
{"with port", "192.0.2.10:5223", "ADVERTISE_IP must be a literal IPv4 address"},
|
||||||
|
{"zero", "0.0.0.0", "ADVERTISE_IP must not be 0.0.0.0"},
|
||||||
|
{"loopback ok", "127.0.0.1", ""},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
env := []string{"MQTT_HOST=mqtt.example.invalid"}
|
||||||
|
if tc.value != "" {
|
||||||
|
env = append(env, "ADVERTISE_IP="+tc.value)
|
||||||
|
}
|
||||||
|
_, err := Load(env)
|
||||||
|
if tc.wantErr == "" {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
|
||||||
|
t.Fatalf("expected error containing %q, got %v", tc.wantErr, err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfigPasswordNotLogged(t *testing.T) {
|
||||||
|
env := append(validEnv(), "MQTT_USERNAME=user", "MQTT_PASSWORD=top-secret-123")
|
||||||
|
cfg, err := Load(env)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.MQTTPassword != "top-secret-123" {
|
||||||
|
t.Fatalf("password not stored")
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, s := range []string{cfg.String(), fmt.Sprintf("%v", cfg), fmt.Sprintf("%+v", cfg), fmt.Sprintf("%#v", cfg)} {
|
||||||
|
if strings.Contains(s, "top-secret") {
|
||||||
|
t.Errorf("redacted output contains password: %q", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
handler := slog.NewTextHandler(&buf, nil)
|
||||||
|
logger := slog.New(handler)
|
||||||
|
logger.Info("test", "config", cfg)
|
||||||
|
if strings.Contains(buf.String(), "top-secret") {
|
||||||
|
t.Errorf("log output contains password: %s", buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultMQTTClientID(t *testing.T) {
|
||||||
|
id := defaultMQTTClientID()
|
||||||
|
if !strings.HasPrefix(id, "n95bridge-") {
|
||||||
|
t.Errorf("defaultMQTTClientID() = %q, want prefix n95bridge-", id)
|
||||||
|
}
|
||||||
|
if !clientIDRe.MatchString(id) {
|
||||||
|
t.Errorf("defaultMQTTClientID() = %q, does not match regex", id)
|
||||||
|
}
|
||||||
|
if len(id) > 64 {
|
||||||
|
t.Errorf("defaultMQTTClientID() length %d > 64", len(id))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeTestCertificate(t *testing.T, path string) {
|
||||||
|
t.Helper()
|
||||||
|
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tmpl := &x509.Certificate{
|
||||||
|
SerialNumber: big.NewInt(1),
|
||||||
|
Subject: pkix.Name{CommonName: "test"},
|
||||||
|
NotBefore: time.Now(),
|
||||||
|
NotAfter: time.Now().Add(time.Hour),
|
||||||
|
}
|
||||||
|
der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
f, err := os.Create(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer f.Close()
|
||||||
|
if err := pem.Encode(f, &pem.Block{Type: "CERTIFICATE", Bytes: der}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,195 @@
|
|||||||
|
package ctl
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CommandTimeout is the deadline on a pending sid/cid correlation, measured
|
||||||
|
// from registration time (milestone plan §10).
|
||||||
|
const CommandTimeout = 30 * time.Second
|
||||||
|
|
||||||
|
// Trace is one correlation event. Phase is "ack" for the stanza receipt and
|
||||||
|
// "result" for the ctl command result or the terminal failure of a command
|
||||||
|
// that expected one. Ret carries the wire ret value or the failure reason.
|
||||||
|
type Trace struct {
|
||||||
|
SID string `json:"sid"`
|
||||||
|
CID string `json:"cid,omitempty"`
|
||||||
|
Command string `json:"command"`
|
||||||
|
Phase string `json:"phase"`
|
||||||
|
Ret string `json:"ret,omitempty"`
|
||||||
|
Errno *string `json:"errno,omitempty"`
|
||||||
|
At time.Time `json:"timestamp"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// TraceSink receives completed traces. A nil sink discards them.
|
||||||
|
type TraceSink func(serial string, trace Trace)
|
||||||
|
|
||||||
|
type record struct {
|
||||||
|
cid string
|
||||||
|
command string
|
||||||
|
expectResult bool
|
||||||
|
sids map[string]struct{}
|
||||||
|
lastSID string
|
||||||
|
registered time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// Correlator tracks outstanding sid and cid waiters. Registering a cid that is
|
||||||
|
// still outstanding attaches the new sid to the single existing record; every
|
||||||
|
// sid may then ack but the cid completes exactly once. A cid is reusable after
|
||||||
|
// completion.
|
||||||
|
type Correlator struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
now func() time.Time
|
||||||
|
bySID map[string]*record
|
||||||
|
byCID map[string]*record
|
||||||
|
records map[*record]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCorrelator creates a Correlator. A nil now uses time.Now.
|
||||||
|
func NewCorrelator(now func() time.Time) *Correlator {
|
||||||
|
if now == nil {
|
||||||
|
now = time.Now
|
||||||
|
}
|
||||||
|
return &Correlator{
|
||||||
|
now: now,
|
||||||
|
bySID: map[string]*record{},
|
||||||
|
byCID: map[string]*record{},
|
||||||
|
records: map[*record]struct{}{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register adds sid as an ack waiter and, when cid is nonempty, a result
|
||||||
|
// waiter. A duplicate registration of a live cid does not allocate a second
|
||||||
|
// waiter; the sid joins the existing record.
|
||||||
|
func (c *Correlator) Register(sid, cid, command string, expectResult bool) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
if cid != "" {
|
||||||
|
if r := c.byCID[cid]; r != nil {
|
||||||
|
r.sids[sid] = struct{}{}
|
||||||
|
r.lastSID = sid
|
||||||
|
c.bySID[sid] = r
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
r := &record{
|
||||||
|
cid: cid,
|
||||||
|
command: command,
|
||||||
|
expectResult: expectResult,
|
||||||
|
sids: map[string]struct{}{sid: {}},
|
||||||
|
lastSID: sid,
|
||||||
|
registered: c.now(),
|
||||||
|
}
|
||||||
|
c.records[r] = struct{}{}
|
||||||
|
c.bySID[sid] = r
|
||||||
|
if cid != "" {
|
||||||
|
c.byCID[cid] = r
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompleteSID completes the stanza ack phase for sid. A sid-only record such
|
||||||
|
// as Move is fully completed by its ack.
|
||||||
|
func (c *Correlator) CompleteSID(sid string) (Trace, bool) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
r := c.bySID[sid]
|
||||||
|
if r == nil {
|
||||||
|
return Trace{}, false
|
||||||
|
}
|
||||||
|
delete(c.bySID, sid)
|
||||||
|
delete(r.sids, sid)
|
||||||
|
tr := Trace{SID: sid, CID: r.cid, Command: r.command, Phase: "ack", At: c.now()}
|
||||||
|
if !r.expectResult && len(r.sids) == 0 {
|
||||||
|
c.remove(r)
|
||||||
|
}
|
||||||
|
return tr, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompleteCID completes the result phase for cid. It returns true only the
|
||||||
|
// first time; the cid is free for reuse afterwards.
|
||||||
|
func (c *Correlator) CompleteCID(cid string) (Trace, bool) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
r := c.byCID[cid]
|
||||||
|
if r == nil {
|
||||||
|
return Trace{}, false
|
||||||
|
}
|
||||||
|
c.remove(r)
|
||||||
|
return Trace{SID: r.lastSID, CID: r.cid, Command: r.command, Phase: "result", At: c.now()}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// FailGeneration emits one terminal trace per outstanding logical command and
|
||||||
|
// clears every waiter. The trace phase is "result" for commands expecting a
|
||||||
|
// ctl result and "ack" for sid-only commands.
|
||||||
|
func (c *Correlator) FailGeneration(reason string) []Trace {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
at := c.now()
|
||||||
|
traces := make([]Trace, 0, len(c.records))
|
||||||
|
for r := range c.records {
|
||||||
|
traces = append(traces, c.terminal(r, reason, at))
|
||||||
|
}
|
||||||
|
c.bySID = map[string]*record{}
|
||||||
|
c.byCID = map[string]*record{}
|
||||||
|
c.records = map[*record]struct{}{}
|
||||||
|
return traces
|
||||||
|
}
|
||||||
|
|
||||||
|
// NextDeadline returns the earliest expiry across outstanding records: each
|
||||||
|
// record's registration time plus CommandTimeout. It reports false when no
|
||||||
|
// records are outstanding.
|
||||||
|
func (c *Correlator) NextDeadline() (time.Time, bool) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
var min time.Time
|
||||||
|
for r := range c.records {
|
||||||
|
d := r.registered.Add(CommandTimeout)
|
||||||
|
if min.IsZero() || d.Before(min) {
|
||||||
|
min = d
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if min.IsZero() {
|
||||||
|
return time.Time{}, false
|
||||||
|
}
|
||||||
|
return min, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expire emits one terminal trace per record registered at least
|
||||||
|
// CommandTimeout before now and removes it.
|
||||||
|
func (c *Correlator) Expire(now time.Time) []Trace {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
var traces []Trace
|
||||||
|
for r := range c.records {
|
||||||
|
if now.Sub(r.registered) >= CommandTimeout {
|
||||||
|
traces = append(traces, c.terminal(r, "timeout", now))
|
||||||
|
c.remove(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return traces
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Correlator) terminal(r *record, reason string, at time.Time) Trace {
|
||||||
|
phase := "ack"
|
||||||
|
if r.expectResult {
|
||||||
|
phase = "result"
|
||||||
|
}
|
||||||
|
return Trace{SID: r.lastSID, CID: r.cid, Command: r.command, Phase: phase, Ret: reason, At: at}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Correlator) remove(r *record) {
|
||||||
|
delete(c.records, r)
|
||||||
|
for sid := range r.sids {
|
||||||
|
delete(c.bySID, sid)
|
||||||
|
}
|
||||||
|
if r.cid != "" {
|
||||||
|
delete(c.byCID, r.cid)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,267 @@
|
|||||||
|
package ctl
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEnvelopeShape(t *testing.T) {
|
||||||
|
stanza, sid, cid, err := Envelope("ctl@ecouser.net/r", "bot@155.ecorobot.net/atom", Outbound{
|
||||||
|
TD: "Clean",
|
||||||
|
Inner: []byte(`<clean type="auto" speed="standard" act="s"/>`),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Envelope: %v", err)
|
||||||
|
}
|
||||||
|
s := string(stanza)
|
||||||
|
if sid == "" || cid == "" {
|
||||||
|
t.Fatalf("sid=%q cid=%q, want both", sid, cid)
|
||||||
|
}
|
||||||
|
if sid == cid {
|
||||||
|
t.Fatalf("sid %q must differ from cid", sid)
|
||||||
|
}
|
||||||
|
if len(cid) != 8 {
|
||||||
|
t.Fatalf("cid %q is not zero-padded to 8 digits", cid)
|
||||||
|
}
|
||||||
|
for _, want := range []string{
|
||||||
|
`<iq id="` + sid + `"`,
|
||||||
|
`type="set"`,
|
||||||
|
`<query xmlns="com:ctl">`,
|
||||||
|
`<ctl td="Clean" id="` + cid + `">`,
|
||||||
|
`<clean type="auto" speed="standard" act="s"/>`,
|
||||||
|
`</ctl></query></iq>`,
|
||||||
|
} {
|
||||||
|
if !strings.Contains(s, want) {
|
||||||
|
t.Errorf("stanza missing %q:\n%s", want, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnvelopeEscapesAttributes(t *testing.T) {
|
||||||
|
stanza, _, _, err := Envelope(`ctl@ecouser.net/"r"`, "bot@x/atom", Outbound{
|
||||||
|
TD: "Clean",
|
||||||
|
Attrs: []Attr{{Name: "speed", Value: `a"b<c&d`}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Envelope: %v", err)
|
||||||
|
}
|
||||||
|
s := string(stanza)
|
||||||
|
if strings.Contains(s, `a"b<c&d`) {
|
||||||
|
t.Fatalf("attribute value not escaped: %s", s)
|
||||||
|
}
|
||||||
|
if !strings.Contains(s, `a"b<c&d`) {
|
||||||
|
t.Fatalf("escaped value missing: %s", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMoveHasNoCtlID(t *testing.T) {
|
||||||
|
stanza, sid, cid, err := Envelope("ctl@ecouser.net/r", "bot@x/atom", Outbound{
|
||||||
|
TD: "Move",
|
||||||
|
OmitCtlID: true,
|
||||||
|
Inner: []byte(`<move action="forward"/>`),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Envelope: %v", err)
|
||||||
|
}
|
||||||
|
if cid != "" {
|
||||||
|
t.Fatalf("Move cid = %q, want empty", cid)
|
||||||
|
}
|
||||||
|
in, err := Parse(stanza)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Parse outbound Move: %v", err)
|
||||||
|
}
|
||||||
|
if in.TD != "Move" {
|
||||||
|
t.Fatalf("td = %q, want Move", in.TD)
|
||||||
|
}
|
||||||
|
if _, has := in.Attrs["id"]; has {
|
||||||
|
t.Fatalf("Move ctl must not carry an id attribute: %s", stanza)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The correlator holds a sid waiter and no cid waiter: the stanza ack
|
||||||
|
// completes the command and no cid can ever complete it.
|
||||||
|
c := NewCorrelator(nil)
|
||||||
|
c.Register(sid, cid, "Move", false)
|
||||||
|
tr, ok := c.CompleteSID(sid)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("Move stanza ack did not complete the sid")
|
||||||
|
}
|
||||||
|
if tr.Phase != "ack" || tr.Command != "Move" {
|
||||||
|
t.Fatalf("trace = %+v, want ack phase for Move", tr)
|
||||||
|
}
|
||||||
|
if _, ok := c.CompleteSID(sid); ok {
|
||||||
|
t.Fatal("second ack completed again")
|
||||||
|
}
|
||||||
|
if traces := c.FailGeneration("connection-lost"); len(traces) != 0 {
|
||||||
|
t.Fatalf("completed Move left waiters: %+v", traces)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseKinds(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
stanza string
|
||||||
|
kind Kind
|
||||||
|
check func(t *testing.T, in Inbound)
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "ack",
|
||||||
|
stanza: `<iq to="c" type="result" id="42"/>`,
|
||||||
|
kind: KindAck,
|
||||||
|
check: func(t *testing.T, in Inbound) {
|
||||||
|
if in.SID != "42" {
|
||||||
|
t.Errorf("sid = %q", in.SID)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "result",
|
||||||
|
stanza: `<iq to="c" type="set" id="9"><query xmlns="com:ctl"><ctl id="00000042" ret="ok" errno=""><battery power="076"/></ctl></query></iq>`,
|
||||||
|
kind: KindResult,
|
||||||
|
check: func(t *testing.T, in Inbound) {
|
||||||
|
if in.CID != "00000042" || in.Ret != "ok" {
|
||||||
|
t.Errorf("cid/ret = %q/%q", in.CID, in.Ret)
|
||||||
|
}
|
||||||
|
if in.Errno == nil || *in.Errno != "" {
|
||||||
|
t.Errorf("errno = %v, want present empty", in.Errno)
|
||||||
|
}
|
||||||
|
if in.BatteryPower != "076" {
|
||||||
|
t.Errorf("battery = %q", in.BatteryPower)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "result errno omitted differs from empty",
|
||||||
|
stanza: `<iq type="set" id="9"><query><ctl id="1" ret="ok"/></query></iq>`,
|
||||||
|
kind: KindResult,
|
||||||
|
check: func(t *testing.T, in Inbound) {
|
||||||
|
if in.Errno != nil {
|
||||||
|
t.Errorf("errno = %q, want nil", *in.Errno)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "push",
|
||||||
|
stanza: `<iq to="c" type="set" id="9"><query xmlns="com:ctl"><ctl td="CleanReport"><clean type="auto" speed="strong" st=" " rsn=" "/></ctl></query></iq>`,
|
||||||
|
kind: KindPush,
|
||||||
|
check: func(t *testing.T, in Inbound) {
|
||||||
|
if in.TD != "CleanReport" {
|
||||||
|
t.Errorf("td = %q", in.TD)
|
||||||
|
}
|
||||||
|
if in.CleanAttrs["type"] != "auto" || in.CleanAttrs["speed"] != "strong" {
|
||||||
|
t.Errorf("clean attrs = %v", in.CleanAttrs)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "bare battery",
|
||||||
|
stanza: `<iq to="c" type="set" id="46"><query xmlns="com:ctl"><battery power="077"/></query></iq>`,
|
||||||
|
kind: KindBattery,
|
||||||
|
check: func(t *testing.T, in Inbound) {
|
||||||
|
if in.BatteryPower != "077" {
|
||||||
|
t.Errorf("battery = %q", in.BatteryPower)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unknown child",
|
||||||
|
stanza: `<iq type="set" id="9"><query xmlns="com:ctl"><frobnicate x="1"/></query></iq>`,
|
||||||
|
kind: KindUnknown,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unqualified elements parse by local name",
|
||||||
|
stanza: `<iq type="set"><query><ctl td="ChargeState"><charge type="going"/></ctl></query></iq>`,
|
||||||
|
kind: KindPush,
|
||||||
|
check: func(t *testing.T, in Inbound) {
|
||||||
|
if in.ChargeAttrs["type"] != "going" {
|
||||||
|
t.Errorf("charge = %v", in.ChargeAttrs)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
in, err := Parse([]byte(tc.stanza))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Parse: %v", err)
|
||||||
|
}
|
||||||
|
if in.Kind != tc.kind {
|
||||||
|
t.Fatalf("kind = %d, want %d", in.Kind, tc.kind)
|
||||||
|
}
|
||||||
|
if tc.check != nil {
|
||||||
|
tc.check(t, in)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseMalformed(t *testing.T) {
|
||||||
|
for _, s := range []string{
|
||||||
|
`<iq type="set"><query><ctl td="CleanReport"`,
|
||||||
|
`not xml at all`,
|
||||||
|
`<iq type="set"><query><ctl td="x"/></query></iq>junk<`,
|
||||||
|
} {
|
||||||
|
if _, err := Parse([]byte(s)); err == nil {
|
||||||
|
t.Errorf("Parse(%q) = nil error, want malformed", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCorrelatorDuplicateLiveCID(t *testing.T) {
|
||||||
|
c := NewCorrelator(nil)
|
||||||
|
c.Register("1", "00000042", "GetCleanState", true)
|
||||||
|
c.Register("2", "00000042", "GetCleanState", true)
|
||||||
|
|
||||||
|
for _, sid := range []string{"1", "2"} {
|
||||||
|
tr, ok := c.CompleteSID(sid)
|
||||||
|
if !ok || tr.Phase != "ack" || tr.CID != "00000042" {
|
||||||
|
t.Fatalf("ack for sid %s = %+v, %v", sid, tr, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tr, ok := c.CompleteCID("00000042")
|
||||||
|
if !ok || tr.Phase != "result" {
|
||||||
|
t.Fatalf("cid completion = %+v, %v", tr, ok)
|
||||||
|
}
|
||||||
|
if _, ok := c.CompleteCID("00000042"); ok {
|
||||||
|
t.Fatal("cid completed twice")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCorrelatorExpireAndFail(t *testing.T) {
|
||||||
|
base := time.Unix(1790194386, 0)
|
||||||
|
now := base
|
||||||
|
c := NewCorrelator(func() time.Time { return now })
|
||||||
|
|
||||||
|
c.Register("1", "00000001", "Clean", true)
|
||||||
|
c.Register("2", "", "Move", false)
|
||||||
|
|
||||||
|
if traces := c.Expire(now.Add(CommandTimeout - time.Second)); len(traces) != 0 {
|
||||||
|
t.Fatalf("early expire produced %+v", traces)
|
||||||
|
}
|
||||||
|
traces := c.Expire(now.Add(CommandTimeout))
|
||||||
|
if len(traces) != 2 {
|
||||||
|
t.Fatalf("expire produced %d traces, want 2", len(traces))
|
||||||
|
}
|
||||||
|
byPhase := map[string]Trace{}
|
||||||
|
for _, tr := range traces {
|
||||||
|
byPhase[tr.Phase] = tr
|
||||||
|
if tr.Ret != "timeout" {
|
||||||
|
t.Errorf("expire ret = %q, want timeout", tr.Ret)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, ok := byPhase["result"]; !ok {
|
||||||
|
t.Error("expectResult command did not expire in result phase")
|
||||||
|
}
|
||||||
|
if _, ok := byPhase["ack"]; !ok {
|
||||||
|
t.Error("sid-only command did not expire in ack phase")
|
||||||
|
}
|
||||||
|
|
||||||
|
c.Register("3", "00000003", "GetSched", true)
|
||||||
|
traces = c.FailGeneration("connection-lost")
|
||||||
|
if len(traces) != 1 || traces[0].Ret != "connection-lost" || traces[0].Phase != "result" {
|
||||||
|
t.Fatalf("FailGeneration = %+v", traces)
|
||||||
|
}
|
||||||
|
if _, ok := c.CompleteSID("3"); ok {
|
||||||
|
t.Fatal("waiter survived FailGeneration")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,90 @@
|
|||||||
|
// Package ctl implements the com:ctl command envelope codec and the sid/cid
|
||||||
|
// correlation tracker from N95-FULL-SPECIFICATION.md §§6 and 14.
|
||||||
|
package ctl
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/xml"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Namespace is the com:ctl query namespace.
|
||||||
|
const Namespace = "com:ctl"
|
||||||
|
|
||||||
|
// Attr is one attribute on the ctl element.
|
||||||
|
type Attr struct {
|
||||||
|
Name string
|
||||||
|
Value string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Outbound describes one ctl command. Inner is an already-built XML fragment
|
||||||
|
// inserted verbatim; Envelope never re-escapes it. OmitCtlID encodes a ctl
|
||||||
|
// element without id, as Move requires.
|
||||||
|
type Outbound struct {
|
||||||
|
TD string
|
||||||
|
Attrs []Attr
|
||||||
|
Inner []byte
|
||||||
|
OmitCtlID bool
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
sidSeq atomic.Uint64
|
||||||
|
cidSeq atomic.Uint64
|
||||||
|
)
|
||||||
|
|
||||||
|
// Envelope wraps out in the §6.1 iq-set envelope. It returns the compact
|
||||||
|
// stanza, the decimal stanza id sid, and the zero-padded 8-digit ctl id cid.
|
||||||
|
// cid is empty when out.OmitCtlID is set. Both ids come from process-wide
|
||||||
|
// counters so they are unique among outstanding operations.
|
||||||
|
func Envelope(from, to string, out Outbound) (stanza []byte, sid, cid string, err error) {
|
||||||
|
if from == "" || to == "" || out.TD == "" {
|
||||||
|
return nil, "", "", fmt.Errorf("ctl: envelope requires from, to and td")
|
||||||
|
}
|
||||||
|
sid = fmt.Sprintf("%d", sidSeq.Add(1))
|
||||||
|
|
||||||
|
var b strings.Builder
|
||||||
|
b.WriteString(`<iq id="`)
|
||||||
|
b.WriteString(escAttr(sid))
|
||||||
|
b.WriteString(`" to="`)
|
||||||
|
b.WriteString(escAttr(to))
|
||||||
|
b.WriteString(`" from="`)
|
||||||
|
b.WriteString(escAttr(from))
|
||||||
|
b.WriteString(`" type="set"><query xmlns="`)
|
||||||
|
b.WriteString(Namespace)
|
||||||
|
b.WriteString(`"><ctl td="`)
|
||||||
|
b.WriteString(escAttr(out.TD))
|
||||||
|
b.WriteString(`"`)
|
||||||
|
if !out.OmitCtlID {
|
||||||
|
cid = fmt.Sprintf("%08d", cidSeq.Add(1))
|
||||||
|
b.WriteString(` id="`)
|
||||||
|
b.WriteString(cid)
|
||||||
|
b.WriteString(`"`)
|
||||||
|
}
|
||||||
|
for _, a := range out.Attrs {
|
||||||
|
b.WriteString(` `)
|
||||||
|
b.WriteString(a.Name)
|
||||||
|
b.WriteString(`="`)
|
||||||
|
b.WriteString(escAttr(a.Value))
|
||||||
|
b.WriteString(`"`)
|
||||||
|
}
|
||||||
|
if len(out.Inner) == 0 {
|
||||||
|
b.WriteString(`/>`)
|
||||||
|
} else {
|
||||||
|
b.WriteString(`>`)
|
||||||
|
b.Write(out.Inner)
|
||||||
|
b.WriteString(`</ctl>`)
|
||||||
|
}
|
||||||
|
b.WriteString(`</query></iq>`)
|
||||||
|
return []byte(b.String()), sid, cid, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func EscapeAttr(s string) string {
|
||||||
|
return escAttr(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
func escAttr(s string) string {
|
||||||
|
var b strings.Builder
|
||||||
|
_ = xml.EscapeText(&b, []byte(s))
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
@@ -0,0 +1,229 @@
|
|||||||
|
package ctl
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/xml"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
var errNoElement = errors.New("ctl: stanza has no root element")
|
||||||
|
|
||||||
|
// Kind classifies an inbound stanza.
|
||||||
|
type Kind uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
KindUnknown Kind = iota
|
||||||
|
// KindAck is an iq result whose id correlates a stanza ack (sid).
|
||||||
|
KindAck
|
||||||
|
// KindResult is an iq set whose ctl carries ret and id (cid).
|
||||||
|
KindResult
|
||||||
|
// KindPush is an unsolicited iq set whose ctl carries td and no ret.
|
||||||
|
KindPush
|
||||||
|
// KindBattery is an iq set whose query holds a bare battery element.
|
||||||
|
KindBattery
|
||||||
|
)
|
||||||
|
|
||||||
|
// Inbound is one parsed inbound stanza. Inner is the raw ctl inner XML.
|
||||||
|
// Errno is a pointer so an omitted errno differs from an empty one.
|
||||||
|
type Inbound struct {
|
||||||
|
Kind Kind
|
||||||
|
SID string
|
||||||
|
CID string
|
||||||
|
TD string
|
||||||
|
Ret string
|
||||||
|
Errno *string
|
||||||
|
Attrs map[string]string
|
||||||
|
Inner []byte
|
||||||
|
BatteryPower string
|
||||||
|
CleanAttrs map[string]string
|
||||||
|
ChargeAttrs map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse classifies one stanza. Elements are matched by local name so both
|
||||||
|
// namespaced and unqualified iq/query forms parse. Well-formed stanzas that
|
||||||
|
// match no known shape return KindUnknown; malformed XML returns an error.
|
||||||
|
func Parse(stanza []byte) (Inbound, error) {
|
||||||
|
var in Inbound
|
||||||
|
dec := xml.NewDecoder(bytes.NewReader(stanza))
|
||||||
|
|
||||||
|
root, err := nextStart(dec)
|
||||||
|
if err != nil {
|
||||||
|
return in, err
|
||||||
|
}
|
||||||
|
if root == nil {
|
||||||
|
return in, errNoElement
|
||||||
|
}
|
||||||
|
if root.Name.Local != "iq" {
|
||||||
|
return in, drain(dec)
|
||||||
|
}
|
||||||
|
iqType := findAttr(root, "type")
|
||||||
|
switch iqType {
|
||||||
|
case "result":
|
||||||
|
if id := findAttr(root, "id"); id != "" {
|
||||||
|
in.Kind = KindAck
|
||||||
|
in.SID = id
|
||||||
|
}
|
||||||
|
return in, drain(dec)
|
||||||
|
case "set":
|
||||||
|
default:
|
||||||
|
return in, drain(dec)
|
||||||
|
}
|
||||||
|
|
||||||
|
query, err := nextStart(dec)
|
||||||
|
if err != nil {
|
||||||
|
return in, err
|
||||||
|
}
|
||||||
|
if query == nil || query.Name.Local != "query" {
|
||||||
|
return in, drain(dec)
|
||||||
|
}
|
||||||
|
|
||||||
|
child, err := nextStart(dec)
|
||||||
|
if err != nil {
|
||||||
|
return in, err
|
||||||
|
}
|
||||||
|
if child == nil {
|
||||||
|
return in, drain(dec)
|
||||||
|
}
|
||||||
|
switch child.Name.Local {
|
||||||
|
case "ctl":
|
||||||
|
if err := parseCtl(dec, child, &in); err != nil {
|
||||||
|
return in, err
|
||||||
|
}
|
||||||
|
case "battery":
|
||||||
|
in.Kind = KindBattery
|
||||||
|
in.BatteryPower = findAttr(child, "power")
|
||||||
|
}
|
||||||
|
return in, drain(dec)
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseCtl(dec *xml.Decoder, start *xml.StartElement, in *Inbound) error {
|
||||||
|
attrs := make(map[string]string, len(start.Attr))
|
||||||
|
for _, a := range start.Attr {
|
||||||
|
attrs[a.Name.Local] = a.Value
|
||||||
|
}
|
||||||
|
in.Attrs = attrs
|
||||||
|
if v, ok := attrs["errno"]; ok {
|
||||||
|
vv := v
|
||||||
|
in.Errno = &vv
|
||||||
|
}
|
||||||
|
|
||||||
|
var inner struct {
|
||||||
|
Raw []byte `xml:",innerxml"`
|
||||||
|
}
|
||||||
|
if err := dec.DecodeElement(&inner, start); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
in.Inner = inner.Raw
|
||||||
|
|
||||||
|
td := attrs["td"]
|
||||||
|
cid := attrs["id"]
|
||||||
|
ret := attrs["ret"]
|
||||||
|
switch {
|
||||||
|
case ret != "" && cid != "":
|
||||||
|
in.Kind = KindResult
|
||||||
|
in.CID = cid
|
||||||
|
in.Ret = ret
|
||||||
|
case td != "" && ret == "":
|
||||||
|
in.Kind = KindPush
|
||||||
|
in.TD = td
|
||||||
|
default:
|
||||||
|
in.Kind = KindUnknown
|
||||||
|
}
|
||||||
|
return parseInner(inner.Raw, in)
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseInner records the attributes of first-level payload elements inside the
|
||||||
|
// ctl inner XML: clean, charge and battery.
|
||||||
|
func parseInner(raw []byte, in *Inbound) error {
|
||||||
|
if len(bytes.TrimSpace(raw)) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
doc := make([]byte, 0, len(raw)+7)
|
||||||
|
doc = append(doc, '<', 'x', '>')
|
||||||
|
doc = append(doc, raw...)
|
||||||
|
doc = append(doc, '<', '/', 'x', '>')
|
||||||
|
dec := xml.NewDecoder(bytes.NewReader(doc))
|
||||||
|
depth := 0
|
||||||
|
for {
|
||||||
|
tok, err := dec.Token()
|
||||||
|
if err == io.EOF {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
switch t := tok.(type) {
|
||||||
|
case xml.StartElement:
|
||||||
|
depth++
|
||||||
|
if depth == 2 {
|
||||||
|
switch t.Name.Local {
|
||||||
|
case "clean":
|
||||||
|
if in.CleanAttrs == nil {
|
||||||
|
in.CleanAttrs = attrsOf(&t)
|
||||||
|
}
|
||||||
|
case "charge":
|
||||||
|
if in.ChargeAttrs == nil {
|
||||||
|
in.ChargeAttrs = attrsOf(&t)
|
||||||
|
}
|
||||||
|
case "battery":
|
||||||
|
if in.BatteryPower == "" {
|
||||||
|
in.BatteryPower = findAttr(&t, "power")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case xml.EndElement:
|
||||||
|
depth--
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func attrsOf(se *xml.StartElement) map[string]string {
|
||||||
|
m := make(map[string]string, len(se.Attr))
|
||||||
|
for _, a := range se.Attr {
|
||||||
|
m[a.Name.Local] = a.Value
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func findAttr(se *xml.StartElement, name string) string {
|
||||||
|
for _, a := range se.Attr {
|
||||||
|
if a.Name.Local == name {
|
||||||
|
return a.Value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// nextStart returns the next start element, skipping text nodes. It returns
|
||||||
|
// nil on EOF or on an end element.
|
||||||
|
func nextStart(dec *xml.Decoder) (*xml.StartElement, error) {
|
||||||
|
for {
|
||||||
|
tok, err := dec.Token()
|
||||||
|
if err == io.EOF {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
switch t := tok.(type) {
|
||||||
|
case xml.StartElement:
|
||||||
|
return &t, nil
|
||||||
|
case xml.EndElement:
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// drain consumes the rest of the stanza so trailing malformed XML is an error.
|
||||||
|
func drain(dec *xml.Decoder) error {
|
||||||
|
for {
|
||||||
|
_, err := dec.Token()
|
||||||
|
if err == io.EOF {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
package ctl
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/xml"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
func escape(s string) string {
|
||||||
|
var b bytes.Buffer
|
||||||
|
xml.EscapeText(&b, []byte(s))
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func AddSchedOutbound(name, on, timeStr, repeat string) Outbound {
|
||||||
|
inner := fmt.Sprintf(`<sched name="%s" on="%s" time="%s" repeat="%s"><ctl td="Clean"><clean type="auto"/></ctl></sched>`,
|
||||||
|
escape(name), escape(on), escape(timeStr), escape(repeat))
|
||||||
|
return Outbound{
|
||||||
|
TD: "AddSched",
|
||||||
|
Inner: []byte(inner),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func ModSchedOutbound(name, on, timeStr, repeat string) Outbound {
|
||||||
|
inner := fmt.Sprintf(`<ModSched name="%s"><sched name="%s" on="%s" time="%s" repeat="%s"><ctl td="Clean"><clean type="auto"/></ctl></sched></ModSched>`,
|
||||||
|
escape(name), escape(name), escape(on), escape(timeStr), escape(repeat))
|
||||||
|
return Outbound{
|
||||||
|
TD: "ModSched",
|
||||||
|
Inner: []byte(inner),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func DelSchedOutbound(name string) Outbound {
|
||||||
|
inner := fmt.Sprintf(`<DelSched name="%s"/>`, escape(name))
|
||||||
|
return Outbound{
|
||||||
|
TD: "DelSched",
|
||||||
|
Inner: []byte(inner),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetSchedOutbound() Outbound {
|
||||||
|
return Outbound{
|
||||||
|
TD: "GetSched",
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
package ha
|
||||||
|
|
||||||
|
import (
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Command maps one MQTT message on a per-robot command topic to a
|
||||||
|
// robot.Command. Payloads are used verbatim — no trimming.
|
||||||
|
func Command(suffix string, payload []byte) robot.Command {
|
||||||
|
s := string(payload)
|
||||||
|
switch suffix {
|
||||||
|
case "command":
|
||||||
|
switch s {
|
||||||
|
case "start", "stop", "return_to_base", "clean_spot", "locate":
|
||||||
|
return robot.Command{Name: s}
|
||||||
|
default:
|
||||||
|
return robot.Command{Reject: s}
|
||||||
|
}
|
||||||
|
case "set_fan_speed":
|
||||||
|
switch s {
|
||||||
|
case "standard", "strong":
|
||||||
|
return robot.Command{
|
||||||
|
Name: "set_fan_speed",
|
||||||
|
Args: map[string]string{"speed": s},
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return robot.Command{Reject: s}
|
||||||
|
}
|
||||||
|
case "send_command":
|
||||||
|
return robot.Command{
|
||||||
|
Name: "send_command",
|
||||||
|
Payload: append([]byte(nil), payload...),
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return robot.Command{Reject: "unsupported"}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
// Package ha produces the Home Assistant MQTT discovery document and maps
|
||||||
|
// inbound MQTT command payloads to robot commands.
|
||||||
|
package ha
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// discoveryDoc is a struct so the compact JSON key order stays fixed.
|
||||||
|
type discoveryDoc struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
UniqueID string `json:"unique_id"`
|
||||||
|
CommandTopic string `json:"command_topic"`
|
||||||
|
SetFanSpeedTopic string `json:"set_fan_speed_topic"`
|
||||||
|
SendCommandTopic string `json:"send_command_topic"`
|
||||||
|
StateTopic string `json:"state_topic"`
|
||||||
|
JSONAttributesTopic string `json:"json_attributes_topic"`
|
||||||
|
AvailabilityTopic string `json:"availability_topic"`
|
||||||
|
PayloadAvailable string `json:"payload_available"`
|
||||||
|
PayloadNotAvailable string `json:"payload_not_available"`
|
||||||
|
FanSpeedList []string `json:"fan_speed_list"`
|
||||||
|
SupportedFeatures []string `json:"supported_features"`
|
||||||
|
Device deviceDiscovery `json:"device"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type deviceDiscovery struct {
|
||||||
|
Identifiers []string `json:"identifiers"`
|
||||||
|
Manufacturer string `json:"manufacturer"`
|
||||||
|
Model string `json:"model"`
|
||||||
|
SerialNumber string `json:"serial_number"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func DiscoveryTopic(cfg config.Config, serial string) string {
|
||||||
|
return cfg.HADiscoveryPrefix + "/vacuum/ecovacs_" + serial + "/config"
|
||||||
|
}
|
||||||
|
|
||||||
|
func Discovery(cfg config.Config, serial string) []byte {
|
||||||
|
root := cfg.MQTTBase + "/" + serial
|
||||||
|
doc := discoveryDoc{
|
||||||
|
Name: "Deebot N95",
|
||||||
|
UniqueID: "ecovacs_" + serial,
|
||||||
|
CommandTopic: root + "/command",
|
||||||
|
SetFanSpeedTopic: root + "/set_fan_speed",
|
||||||
|
SendCommandTopic: root + "/send_command",
|
||||||
|
StateTopic: root + "/state",
|
||||||
|
JSONAttributesTopic: root + "/json_attributes",
|
||||||
|
AvailabilityTopic: root + "/availability",
|
||||||
|
PayloadAvailable: "online",
|
||||||
|
PayloadNotAvailable: "offline",
|
||||||
|
FanSpeedList: []string{"standard", "strong"},
|
||||||
|
SupportedFeatures: []string{
|
||||||
|
"start", "stop", "return_home", "status", "locate",
|
||||||
|
"clean_spot", "fan_speed", "send_command",
|
||||||
|
},
|
||||||
|
Device: deviceDiscovery{
|
||||||
|
Identifiers: []string{serial},
|
||||||
|
Manufacturer: "Ecovacs",
|
||||||
|
Model: "Deebot N95 (wukong/155)",
|
||||||
|
SerialNumber: serial,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
b, _ := json.Marshal(doc)
|
||||||
|
return b
|
||||||
|
}
|
||||||
@@ -0,0 +1,156 @@
|
|||||||
|
package ha
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
)
|
||||||
|
|
||||||
|
func testConfig() config.Config {
|
||||||
|
return config.Config{
|
||||||
|
MQTTBase: "ecovacs",
|
||||||
|
HADiscoveryPrefix: "homeassistant",
|
||||||
|
MQTTClientID: "n95-test",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func discoveryMap(t *testing.T, cfg config.Config, serial string) map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
var doc map[string]any
|
||||||
|
if err := json.Unmarshal(Discovery(cfg, serial), &doc); err != nil {
|
||||||
|
t.Fatalf("discovery JSON invalid: %v", err)
|
||||||
|
}
|
||||||
|
return doc
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoveryJSONOmitsPlatformAndBattery(t *testing.T) {
|
||||||
|
cfg := testConfig()
|
||||||
|
serial := "SER123"
|
||||||
|
payload := Discovery(cfg, serial)
|
||||||
|
|
||||||
|
want := `{"name":"Deebot N95","unique_id":"ecovacs_SER123",` +
|
||||||
|
`"command_topic":"ecovacs/SER123/command",` +
|
||||||
|
`"set_fan_speed_topic":"ecovacs/SER123/set_fan_speed",` +
|
||||||
|
`"send_command_topic":"ecovacs/SER123/send_command",` +
|
||||||
|
`"state_topic":"ecovacs/SER123/state",` +
|
||||||
|
`"json_attributes_topic":"ecovacs/SER123/json_attributes",` +
|
||||||
|
`"availability_topic":"ecovacs/SER123/availability",` +
|
||||||
|
`"payload_available":"online","payload_not_available":"offline",` +
|
||||||
|
`"fan_speed_list":["standard","strong"],` +
|
||||||
|
`"supported_features":["start","stop","return_home","status","locate","clean_spot","fan_speed","send_command"],` +
|
||||||
|
`"device":{"identifiers":["SER123"],"manufacturer":"Ecovacs","model":"Deebot N95 (wukong/155)","serial_number":"SER123"}}`
|
||||||
|
if string(payload) != want {
|
||||||
|
t.Fatalf("discovery payload mismatch\ngot: %s\nwant: %s", payload, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
doc := discoveryMap(t, cfg, serial)
|
||||||
|
if _, ok := doc["platform"]; ok {
|
||||||
|
t.Fatal("discovery must not contain platform")
|
||||||
|
}
|
||||||
|
device, ok := doc["device"].(map[string]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("device must be an object")
|
||||||
|
}
|
||||||
|
if _, ok := device["connections"]; ok {
|
||||||
|
t.Fatal("device must not contain connections")
|
||||||
|
}
|
||||||
|
for _, banned := range []string{"battery", "pause", "clean_segments"} {
|
||||||
|
if _, ok := doc[banned]; ok {
|
||||||
|
t.Fatalf("discovery must not contain %s", banned)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoveryRetainedPrefix(t *testing.T) {
|
||||||
|
cfg := testConfig()
|
||||||
|
cfg.HADiscoveryPrefix = "ha2"
|
||||||
|
cfg.MQTTBase = "vac"
|
||||||
|
serial := "ABC999"
|
||||||
|
|
||||||
|
if got, want := DiscoveryTopic(cfg, serial), "ha2/vacuum/ecovacs_ABC999/config"; got != want {
|
||||||
|
t.Fatalf("DiscoveryTopic = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
doc := discoveryMap(t, cfg, serial)
|
||||||
|
if got := doc["unique_id"]; got != "ecovacs_ABC999" {
|
||||||
|
t.Fatalf("unique_id = %v; object id must keep ecovacs_{serial}", got)
|
||||||
|
}
|
||||||
|
if got := doc["command_topic"]; got != "vac/ABC999/command" {
|
||||||
|
t.Fatalf("command_topic = %v; must use MQTT_BASE", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoveryOmitsUnverifiedFeatures(t *testing.T) {
|
||||||
|
doc := discoveryMap(t, testConfig(), "SER1")
|
||||||
|
features, ok := doc["supported_features"].([]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("supported_features must be an array")
|
||||||
|
}
|
||||||
|
want := []string{"start", "stop", "return_home", "status", "locate", "clean_spot", "fan_speed", "send_command"}
|
||||||
|
var got []string
|
||||||
|
for _, f := range features {
|
||||||
|
got = append(got, f.(string))
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Fatalf("supported_features = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
for _, banned := range []string{"battery", "pause", "clean_segments"} {
|
||||||
|
for _, f := range got {
|
||||||
|
if f == banned {
|
||||||
|
t.Fatalf("unverified feature %q advertised", banned)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoveryFanSpeedList(t *testing.T) {
|
||||||
|
doc := discoveryMap(t, testConfig(), "SER1")
|
||||||
|
list, ok := doc["fan_speed_list"].([]any)
|
||||||
|
if !ok || len(list) != 2 || list[0] != "standard" || list[1] != "strong" {
|
||||||
|
t.Fatalf("fan_speed_list = %v, want [standard strong]", doc["fan_speed_list"])
|
||||||
|
}
|
||||||
|
if got := robot.NewSnapshot().State.FanSpeed; got != "standard" {
|
||||||
|
t.Fatalf("initial fan_speed = %q, want standard", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCommandMappingTable(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
suffix, payload string
|
||||||
|
want robot.Command
|
||||||
|
}{
|
||||||
|
{"command", "start", robot.Command{Name: "start"}},
|
||||||
|
{"command", "stop", robot.Command{Name: "stop"}},
|
||||||
|
{"command", "return_to_base", robot.Command{Name: "return_to_base"}},
|
||||||
|
{"command", "clean_spot", robot.Command{Name: "clean_spot"}},
|
||||||
|
{"command", "locate", robot.Command{Name: "locate"}},
|
||||||
|
{"command", "pause", robot.Command{Reject: "pause"}},
|
||||||
|
{"command", "garbage", robot.Command{Reject: "garbage"}},
|
||||||
|
{"set_fan_speed", "standard", robot.Command{Name: "set_fan_speed", Args: map[string]string{"speed": "standard"}}},
|
||||||
|
{"set_fan_speed", "strong", robot.Command{Name: "set_fan_speed", Args: map[string]string{"speed": "strong"}}},
|
||||||
|
{"set_fan_speed", "turbo", robot.Command{Reject: "turbo"}},
|
||||||
|
{"unknown_suffix", "x", robot.Command{Reject: "unsupported"}},
|
||||||
|
} {
|
||||||
|
got := Command(tc.suffix, []byte(tc.payload))
|
||||||
|
if !reflect.DeepEqual(got, tc.want) {
|
||||||
|
t.Errorf("Command(%q, %q) = %+v, want %+v", tc.suffix, tc.payload, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendCommandCopiesPayload(t *testing.T) {
|
||||||
|
in := []byte(`{"command":"move"}`)
|
||||||
|
cmd := Command("send_command", in)
|
||||||
|
if cmd.Name != "send_command" {
|
||||||
|
t.Fatalf("Name = %q, want send_command", cmd.Name)
|
||||||
|
}
|
||||||
|
if string(cmd.Payload) != string(in) {
|
||||||
|
t.Fatalf("Payload = %q, want %q", cmd.Payload, in)
|
||||||
|
}
|
||||||
|
in[0] = 'X'
|
||||||
|
if cmd.Payload[0] == 'X' {
|
||||||
|
t.Fatal("Payload aliases the input slice; must be a copy")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,189 @@
|
|||||||
|
package ha
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"encoding/xml"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
)
|
||||||
|
|
||||||
|
type sendCommandJSON struct {
|
||||||
|
Command string `json:"command"`
|
||||||
|
CleanType string `json:"clean_type"`
|
||||||
|
Action string `json:"action"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
On string `json:"on"`
|
||||||
|
Time string `json:"time"`
|
||||||
|
Repeat string `json:"repeat"`
|
||||||
|
XML string `json:"xml"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func SendCommandHook(cfg config.Config) func(ctx context.Context, serial string, payload []byte) (robot.Command, error) {
|
||||||
|
return func(ctx context.Context, serial string, payload []byte) (robot.Command, error) {
|
||||||
|
if len(payload) > 4096 {
|
||||||
|
return robot.Command{}, errors.New("rejected:size")
|
||||||
|
}
|
||||||
|
|
||||||
|
var req sendCommandJSON
|
||||||
|
if err := json.Unmarshal(payload, &req); err != nil {
|
||||||
|
return robot.Command{}, errors.New("rejected:payload")
|
||||||
|
}
|
||||||
|
|
||||||
|
if req.Command == "" {
|
||||||
|
return robot.Command{}, errors.New("rejected:payload")
|
||||||
|
}
|
||||||
|
|
||||||
|
switch req.Command {
|
||||||
|
case "clean":
|
||||||
|
if req.CleanType != "auto" && req.CleanType != "border" && req.CleanType != "spot" && req.CleanType != "singleRoom" {
|
||||||
|
return robot.Command{}, fmt.Errorf("rejected:clean-type")
|
||||||
|
}
|
||||||
|
return robot.Command{Name: "clean", Args: map[string]string{"type": req.CleanType}}, nil
|
||||||
|
|
||||||
|
case "move":
|
||||||
|
if req.Action == "backward" {
|
||||||
|
return robot.Command{}, errors.New("rejected:backward")
|
||||||
|
}
|
||||||
|
if req.Action != "forward" && req.Action != "SpinLeft" && req.Action != "SpinRight" && req.Action != "TurnAround" && req.Action != "stop" {
|
||||||
|
return robot.Command{}, errors.New("rejected:move-action")
|
||||||
|
}
|
||||||
|
return robot.Command{Name: "move", Args: map[string]string{"action": req.Action}}, nil
|
||||||
|
|
||||||
|
case "cancel_return":
|
||||||
|
return robot.Command{Name: "cancel_return"}, nil
|
||||||
|
|
||||||
|
case "set_time":
|
||||||
|
return robot.Command{Name: "set_time"}, nil
|
||||||
|
|
||||||
|
case "get_status":
|
||||||
|
return robot.Command{Name: "get_status"}, nil
|
||||||
|
|
||||||
|
case "get_lifespan":
|
||||||
|
return robot.Command{Name: "get_lifespan"}, nil
|
||||||
|
|
||||||
|
case "add_sched", "mod_sched":
|
||||||
|
if req.Name == "" || len(req.Name) > 64 {
|
||||||
|
return robot.Command{}, errors.New("rejected:name")
|
||||||
|
}
|
||||||
|
if req.On != "1" && req.On != "0" && req.On != "true" && req.On != "false" {
|
||||||
|
return robot.Command{}, errors.New("rejected:on")
|
||||||
|
}
|
||||||
|
if req.On == "true" {
|
||||||
|
req.On = "1"
|
||||||
|
} else if req.On == "false" {
|
||||||
|
req.On = "0"
|
||||||
|
}
|
||||||
|
if req.CleanType != "auto" {
|
||||||
|
return robot.Command{}, errors.New("rejected:schedule-type")
|
||||||
|
}
|
||||||
|
if len(req.Repeat) != 7 {
|
||||||
|
return robot.Command{}, errors.New("rejected:repeat")
|
||||||
|
}
|
||||||
|
for _, c := range req.Repeat {
|
||||||
|
if c != '0' && c != '1' {
|
||||||
|
return robot.Command{}, errors.New("rejected:repeat")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(req.Time) != 5 || req.Time[2] != ':' {
|
||||||
|
return robot.Command{}, errors.New("rejected:time")
|
||||||
|
}
|
||||||
|
return robot.Command{
|
||||||
|
Name: req.Command,
|
||||||
|
Args: map[string]string{
|
||||||
|
"name": req.Name,
|
||||||
|
"on": req.On,
|
||||||
|
"time": req.Time,
|
||||||
|
"repeat": req.Repeat,
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
|
||||||
|
case "del_sched":
|
||||||
|
return robot.Command{Name: "del_sched", Args: map[string]string{"name": req.Name}}, nil
|
||||||
|
|
||||||
|
case "get_sched":
|
||||||
|
return robot.Command{Name: "get_sched"}, nil
|
||||||
|
|
||||||
|
case "resume":
|
||||||
|
return robot.Command{}, errors.New("rejected:resume")
|
||||||
|
|
||||||
|
case "raw":
|
||||||
|
if !cfg.RawCommands {
|
||||||
|
return robot.Command{}, errors.New("raw-disabled")
|
||||||
|
}
|
||||||
|
if len(req.XML) > 4096 {
|
||||||
|
return robot.Command{}, errors.New("rejected:size")
|
||||||
|
}
|
||||||
|
dec := xml.NewDecoder(strings.NewReader(req.XML))
|
||||||
|
tok, err := dec.Token()
|
||||||
|
if err != nil {
|
||||||
|
return robot.Command{}, errors.New("rejected:raw-xml")
|
||||||
|
}
|
||||||
|
var start xml.StartElement
|
||||||
|
var ok bool
|
||||||
|
if start, ok = tok.(xml.StartElement); !ok {
|
||||||
|
return robot.Command{}, errors.New("rejected:raw-xml")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reject forbidden tags
|
||||||
|
switch start.Name.Local {
|
||||||
|
case "stream", "auth", "iq", "presence", "message", "starttls", "bind", "session":
|
||||||
|
return robot.Command{}, errors.New("rejected:raw-xml-root")
|
||||||
|
}
|
||||||
|
|
||||||
|
if start.Name.Local != "ctl" {
|
||||||
|
return robot.Command{}, errors.New("rejected:raw-xml-root")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Capture inner XML
|
||||||
|
var innerBuf bytes.Buffer
|
||||||
|
enc := xml.NewEncoder(&innerBuf)
|
||||||
|
|
||||||
|
var idAttr string
|
||||||
|
var tdAttr string
|
||||||
|
|
||||||
|
for _, attr := range start.Attr {
|
||||||
|
if attr.Name.Local == "id" {
|
||||||
|
idAttr = attr.Value
|
||||||
|
} else if attr.Name.Local == "td" {
|
||||||
|
tdAttr = attr.Value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if idAttr == "" && tdAttr != "Move" {
|
||||||
|
return robot.Command{}, errors.New("rejected:raw-id")
|
||||||
|
}
|
||||||
|
|
||||||
|
for {
|
||||||
|
t, err := dec.Token()
|
||||||
|
if err != nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if end, ok := t.(xml.EndElement); ok && end.Name.Local == "ctl" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err := enc.EncodeToken(t); err != nil {
|
||||||
|
return robot.Command{}, errors.New("rejected:raw-xml-inner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
enc.Flush()
|
||||||
|
|
||||||
|
return robot.Command{
|
||||||
|
Name: "raw",
|
||||||
|
Args: map[string]string{
|
||||||
|
"td": tdAttr,
|
||||||
|
"id": idAttr,
|
||||||
|
},
|
||||||
|
Payload: innerBuf.Bytes(),
|
||||||
|
}, nil
|
||||||
|
|
||||||
|
default:
|
||||||
|
return robot.Command{}, errors.New("rejected:" + req.Command)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
package ha_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ha"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSendCommandExtensions(t *testing.T) {
|
||||||
|
cfg := config.Config{RawCommands: false}
|
||||||
|
hook := ha.SendCommandHook(cfg)
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
payload string
|
||||||
|
want string
|
||||||
|
err string
|
||||||
|
}{
|
||||||
|
{`{"command":"clean","clean_type":"border"}`, "clean", ""},
|
||||||
|
{`{"command":"move","action":"SpinLeft"}`, "move", ""},
|
||||||
|
{`{"command":"cancel_return"}`, "cancel_return", ""},
|
||||||
|
{`{"command":"set_time"}`, "set_time", ""},
|
||||||
|
{`{"command":"get_status"}`, "get_status", ""},
|
||||||
|
{`{"command":"get_lifespan"}`, "get_lifespan", ""},
|
||||||
|
{`{"command":"add_sched","name":"123","on":"1","time":"21:30","repeat":"0111000","clean_type":"auto"}`, "add_sched", ""},
|
||||||
|
{`{"command":"resume"}`, "", "rejected:resume"},
|
||||||
|
{`{"command":"move","action":"backward"}`, "", "rejected:backward"},
|
||||||
|
{`{"command":"add_sched","name":"123","on":"1","time":"21:30","repeat":"999","clean_type":"auto"}`, "", "rejected:repeat"},
|
||||||
|
{`{"command":"add_sched","name":"123","on":"1","time":"2130","repeat":"0111000","clean_type":"auto"}`, "", "rejected:time"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
cmd, err := hook(context.Background(), "SER1", []byte(tc.payload))
|
||||||
|
if tc.err != "" {
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.err) {
|
||||||
|
t.Errorf("payload %s: want err %s, got %v", tc.payload, tc.err, err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("payload %s: unexpected err: %v", tc.payload, err)
|
||||||
|
} else if cmd.Name != tc.want {
|
||||||
|
t.Errorf("payload %s: want name %s, got %s", tc.payload, tc.want, cmd.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRawCommandsDefaultOff(t *testing.T) {
|
||||||
|
cfgOff := config.Config{RawCommands: false}
|
||||||
|
hookOff := ha.SendCommandHook(cfgOff)
|
||||||
|
_, err := hookOff(context.Background(), "SER1", []byte(`{"command":"raw","xml":"<ctl id='1'/>"}`))
|
||||||
|
if err == nil || err.Error() != "raw-disabled" {
|
||||||
|
t.Errorf("expected raw-disabled, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfgOn := config.Config{RawCommands: true}
|
||||||
|
hookOn := ha.SendCommandHook(cfgOn)
|
||||||
|
cmd, err := hookOn(context.Background(), "SER1", []byte(`{"command":"raw","xml":"<ctl id='1'/>"}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if cmd.Name != "raw" {
|
||||||
|
t.Errorf("expected raw command, got %s", cmd.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = hookOn(context.Background(), "SER1", []byte(`{"command":"raw","xml":"<stream/>"}`))
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "rejected:") {
|
||||||
|
t.Errorf("expected stream root to be rejected, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
package httpx
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ServeFirmware runs the firmware 404 listener until ctx is cancelled.
|
||||||
|
func ServeFirmware(ctx context.Context, cfg config.Config) error {
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.HandleFunc("/", firmwareHandler)
|
||||||
|
mux.HandleFunc("/products/wukong/class/155/firmware/latest.json", firmwareHandler)
|
||||||
|
|
||||||
|
addr := net.JoinHostPort(cfg.BindAddress, fmt.Sprintf("%d", cfg.PortFirmware))
|
||||||
|
ln, err := net.Listen("tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s: %w", addr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
srv := &http.Server{
|
||||||
|
Handler: mux,
|
||||||
|
ReadHeaderTimeout: 5 * time.Second,
|
||||||
|
ReadTimeout: 10 * time.Second,
|
||||||
|
WriteTimeout: 10 * time.Second,
|
||||||
|
}
|
||||||
|
srv.SetKeepAlivesEnabled(false)
|
||||||
|
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
errCh <- srv.Serve(ln)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
_ = srv.Shutdown(shutdownCtx)
|
||||||
|
return nil
|
||||||
|
case err := <-errCh:
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func firmwareHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||||
|
w.Header().Set("Content-Length", "9")
|
||||||
|
w.Header().Set("Connection", "close")
|
||||||
|
w.WriteHeader(http.StatusNotFound)
|
||||||
|
_, _ = w.Write([]byte("Not Found"))
|
||||||
|
}
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
package httpx
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ServeHealth runs the GET /healthz listener until ctx is cancelled.
|
||||||
|
func ServeHealth(ctx context.Context, cfg config.Config) error {
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.HandleFunc("/healthz", healthHandler)
|
||||||
|
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusNotFound)
|
||||||
|
})
|
||||||
|
|
||||||
|
addr := net.JoinHostPort(cfg.BindAddress, fmt.Sprintf("%d", cfg.HealthPort))
|
||||||
|
ln, err := net.Listen("tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s: %w", addr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
srv := &http.Server{
|
||||||
|
Handler: mux,
|
||||||
|
ReadHeaderTimeout: 5 * time.Second,
|
||||||
|
ReadTimeout: 10 * time.Second,
|
||||||
|
WriteTimeout: 10 * time.Second,
|
||||||
|
}
|
||||||
|
srv.SetKeepAlivesEnabled(false)
|
||||||
|
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
errCh <- srv.Serve(ln)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
_ = srv.Shutdown(shutdownCtx)
|
||||||
|
return nil
|
||||||
|
case err := <-errCh:
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func healthHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte("ok\n"))
|
||||||
|
}
|
||||||
@@ -0,0 +1,323 @@
|
|||||||
|
package httpx
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func baseConfig(extra ...string) config.Config {
|
||||||
|
env := []string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
}
|
||||||
|
env = append(env, extra...)
|
||||||
|
cfg, err := config.Load(env)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func startServer(t *testing.T, fn func(context.Context, config.Config) error, cfg config.Config) context.CancelFunc {
|
||||||
|
t.Helper()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(done)
|
||||||
|
_ = fn(ctx, cfg)
|
||||||
|
}()
|
||||||
|
t.Cleanup(func() {
|
||||||
|
cancel()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("server did not stop")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
return cancel
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupHTTP10WithoutHost(t *testing.T) {
|
||||||
|
cfg := baseConfig("PORT_LOOKUP=18007")
|
||||||
|
startServer(t, ServeLookup, cfg)
|
||||||
|
|
||||||
|
conn, err := net.Dial("tcp", "127.0.0.1:18007")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial: %v", err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
body := `{"todo":"FindBest","service":"EcoMsgNew"}`
|
||||||
|
fmt.Fprintf(conn, "POST /lookup.do HTTP/1.0\r\nContent-Length: %d\r\n\r\n%s", len(body), body)
|
||||||
|
|
||||||
|
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read response: %v", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want 200", resp.StatusCode)
|
||||||
|
}
|
||||||
|
if resp.Header.Get("Content-Type") != "application/json; charset=utf-8" {
|
||||||
|
t.Errorf("content-type = %q", resp.Header.Get("Content-Type"))
|
||||||
|
}
|
||||||
|
b, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read body: %v", err)
|
||||||
|
}
|
||||||
|
want := `{"result":"ok","ip":"192.0.2.10","port":5223}`
|
||||||
|
if string(b) != want {
|
||||||
|
t.Errorf("body = %q, want %q", string(b), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupCompactNumericPort(t *testing.T) {
|
||||||
|
cfg := baseConfig(
|
||||||
|
"PORT_XMPP=18223",
|
||||||
|
"PORT_FIRMWARE=18005",
|
||||||
|
"PORT_LOOKUP=18007",
|
||||||
|
)
|
||||||
|
startServer(t, ServeLookup, cfg)
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
service string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"EcoMsgNew", `{"result":"ok","ip":"192.0.2.10","port":18223}`},
|
||||||
|
{"EcoUpdate", `{"result":"ok","ip":"192.0.2.10","port":18005}`},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
conn, err := net.Dial("tcp", "127.0.0.1:18007")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial: %v", err)
|
||||||
|
}
|
||||||
|
body := fmt.Sprintf(`{"todo":"FindBest","service":"%s"}`, tc.service)
|
||||||
|
fmt.Fprintf(conn, "POST /lookup.do HTTP/1.0\r\nContent-Length: %d\r\n\r\n%s", len(body), body)
|
||||||
|
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
t.Fatalf("read response: %v", err)
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
conn.Close()
|
||||||
|
if string(b) != tc.want {
|
||||||
|
t.Errorf("service %s: body = %q, want %q", tc.service, string(b), tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParallelLookupIndependentSockets(t *testing.T) {
|
||||||
|
cfg := baseConfig(
|
||||||
|
"PORT_XMPP=18223",
|
||||||
|
"PORT_FIRMWARE=18005",
|
||||||
|
"PORT_LOOKUP=18007",
|
||||||
|
)
|
||||||
|
startServer(t, ServeLookup, cfg)
|
||||||
|
|
||||||
|
ch := make(chan string, 2)
|
||||||
|
dialAndRead := func(service, want string) {
|
||||||
|
conn, err := net.Dial("tcp", "127.0.0.1:18007")
|
||||||
|
if err != nil {
|
||||||
|
ch <- fmt.Sprintf("dial %s: %v", service, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
body := fmt.Sprintf(`{"todo":"FindBest","service":"%s"}`, service)
|
||||||
|
fmt.Fprintf(conn, "POST /lookup.do HTTP/1.0\r\nContent-Length: %d\r\n\r\n%s", len(body), body)
|
||||||
|
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
ch <- fmt.Sprintf("read %s: %v", service, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
conn.Close()
|
||||||
|
if string(b) != want {
|
||||||
|
ch <- fmt.Sprintf("%s: got %q want %q", service, string(b), want)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ch <- ""
|
||||||
|
}
|
||||||
|
|
||||||
|
go dialAndRead("EcoMsgNew", `{"result":"ok","ip":"192.0.2.10","port":18223}`)
|
||||||
|
go dialAndRead("EcoUpdate", `{"result":"ok","ip":"192.0.2.10","port":18005}`)
|
||||||
|
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
if msg := <-ch; msg != "" {
|
||||||
|
t.Error(msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupFailAndRST(t *testing.T) {
|
||||||
|
cfg := baseConfig("PORT_LOOKUP=18007")
|
||||||
|
startServer(t, ServeLookup, cfg)
|
||||||
|
|
||||||
|
conn, err := net.Dial("tcp", "127.0.0.1:18007")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial: %v", err)
|
||||||
|
}
|
||||||
|
body := `{"todo":"FindBest","service":"UnknownService"}`
|
||||||
|
fmt.Fprintf(conn, "POST /lookup.do HTTP/1.0\r\nContent-Length: %d\r\n\r\n%s", len(body), body)
|
||||||
|
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read response: %v", err)
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
conn.Close()
|
||||||
|
if string(b) != `{"result":"fail"}` {
|
||||||
|
t.Errorf("unknown service body = %q, want fail", string(b))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Force an RST on a second connection by setting SO_LINGER to zero.
|
||||||
|
conn2, err := net.Dial("tcp", "127.0.0.1:18007")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dial: %v", err)
|
||||||
|
}
|
||||||
|
body = `{"todo":"FindBest","service":"EcoMsgNew"}`
|
||||||
|
fmt.Fprintf(conn2, "POST /lookup.do HTTP/1.0\r\nContent-Length: %d\r\n\r\n%s", len(body), body)
|
||||||
|
// Give the server time to process and close; then RST this side.
|
||||||
|
if tcp, ok := conn2.(*net.TCPConn); ok {
|
||||||
|
tcp.SetLinger(0)
|
||||||
|
}
|
||||||
|
conn2.Close()
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
// The server must still answer a fresh request.
|
||||||
|
resp, err = http.Post("http://127.0.0.1:18007/lookup.do", "application/json", strings.NewReader(body))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("post after rst: %v", err)
|
||||||
|
}
|
||||||
|
b, _ = io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
want := `{"result":"ok","ip":"192.0.2.10","port":5223}`
|
||||||
|
if string(b) != want {
|
||||||
|
t.Errorf("body after rst = %q, want %q", string(b), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFirmware404ExactBytes(t *testing.T) {
|
||||||
|
cfg := baseConfig(
|
||||||
|
"PORT_FIRMWARE=18005",
|
||||||
|
"PORT_LOOKUP=18007",
|
||||||
|
)
|
||||||
|
startServer(t, ServeFirmware, cfg)
|
||||||
|
|
||||||
|
for _, path := range []string{
|
||||||
|
"/products/wukong/class/155/firmware/latest.json",
|
||||||
|
"/other",
|
||||||
|
} {
|
||||||
|
resp, err := http.Get("http://127.0.0.1:18005" + path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get %s: %v", path, err)
|
||||||
|
}
|
||||||
|
if resp.StatusCode != http.StatusNotFound {
|
||||||
|
t.Errorf("path %s status = %d, want 404", path, resp.StatusCode)
|
||||||
|
}
|
||||||
|
if ct := resp.Header.Get("Content-Type"); ct != "text/plain; charset=utf-8" {
|
||||||
|
t.Errorf("path %s content-type = %q", path, ct)
|
||||||
|
}
|
||||||
|
if cl := resp.Header.Get("Content-Length"); cl != "9" {
|
||||||
|
t.Errorf("path %s content-length = %q, want 9", path, cl)
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
if string(b) != "Not Found" {
|
||||||
|
t.Errorf("path %s body = %q, want %q", path, string(b), "Not Found")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAdvertiseIPIsNotBindAddress(t *testing.T) {
|
||||||
|
cfg := baseConfig(
|
||||||
|
"BIND_ADDRESS=127.0.0.1",
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"PORT_LOOKUP=18007",
|
||||||
|
)
|
||||||
|
startServer(t, ServeLookup, cfg)
|
||||||
|
|
||||||
|
body := `{"todo":"FindBest","service":"EcoMsgNew"}`
|
||||||
|
resp, err := http.Post("http://127.0.0.1:18007/lookup.do", "application/json", strings.NewReader(body))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("post: %v", err)
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
want := `{"result":"ok","ip":"192.0.2.10","port":5223}`
|
||||||
|
if string(b) != want {
|
||||||
|
t.Errorf("body = %q, want %q", string(b), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthzSeparatePort(t *testing.T) {
|
||||||
|
cfg := baseConfig(
|
||||||
|
"PORT_FIRMWARE=18005",
|
||||||
|
"PORT_LOOKUP=18007",
|
||||||
|
"HEALTH_PORT=18080",
|
||||||
|
)
|
||||||
|
startServer(t, ServeFirmware, cfg)
|
||||||
|
startServer(t, ServeLookup, cfg)
|
||||||
|
startServer(t, ServeHealth, cfg)
|
||||||
|
|
||||||
|
resp, err := http.Get("http://127.0.0.1:18080/healthz")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get healthz: %v", err)
|
||||||
|
}
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
t.Errorf("healthz status = %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
if ct := resp.Header.Get("Content-Type"); ct != "text/plain; charset=utf-8" {
|
||||||
|
t.Errorf("healthz content-type = %q", ct)
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
if string(b) != "ok\n" {
|
||||||
|
t.Errorf("healthz body = %q, want ok\\n", string(b))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, url := range []string{
|
||||||
|
"http://127.0.0.1:18005/healthz",
|
||||||
|
"http://127.0.0.1:18007/healthz",
|
||||||
|
} {
|
||||||
|
resp, err := http.Get(url)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get %s: %v", url, err)
|
||||||
|
}
|
||||||
|
if resp.StatusCode != http.StatusNotFound {
|
||||||
|
t.Errorf("%s status = %d, want 404", url, resp.StatusCode)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthzUnknownPath(t *testing.T) {
|
||||||
|
cfg := baseConfig("HEALTH_PORT=18080")
|
||||||
|
startServer(t, ServeHealth, cfg)
|
||||||
|
|
||||||
|
resp, err := http.Get("http://127.0.0.1:18080/other")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get: %v", err)
|
||||||
|
}
|
||||||
|
if resp.StatusCode != http.StatusNotFound {
|
||||||
|
t.Errorf("status = %d, want 404", resp.StatusCode)
|
||||||
|
}
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
if len(b) != 0 {
|
||||||
|
t.Errorf("body = %q, want empty", string(b))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,106 @@
|
|||||||
|
package httpx
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ServeLookup runs the POST /lookup.do listener until ctx is cancelled.
|
||||||
|
func ServeLookup(ctx context.Context, cfg config.Config) error {
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.HandleFunc("/lookup.do", newLookupHandler(cfg))
|
||||||
|
|
||||||
|
addr := net.JoinHostPort(cfg.BindAddress, fmt.Sprintf("%d", cfg.PortLookup))
|
||||||
|
ln, err := net.Listen("tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s: %w", addr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
srv := &http.Server{
|
||||||
|
Handler: mux,
|
||||||
|
ReadHeaderTimeout: 5 * time.Second,
|
||||||
|
ReadTimeout: 10 * time.Second,
|
||||||
|
WriteTimeout: 10 * time.Second,
|
||||||
|
}
|
||||||
|
srv.SetKeepAlivesEnabled(false)
|
||||||
|
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
errCh <- srv.Serve(ln)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
_ = srv.Shutdown(shutdownCtx)
|
||||||
|
return nil
|
||||||
|
case err := <-errCh:
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newLookupHandler(cfg config.Config) http.HandlerFunc {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||||
|
w.Header().Set("Connection", "close")
|
||||||
|
|
||||||
|
r.Body = http.MaxBytesReader(w, r.Body, 64*1024)
|
||||||
|
defer r.Body.Close()
|
||||||
|
|
||||||
|
var req lookupRequest
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
slog.Debug("lookup request decode failed", "err", err)
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if req.Todo != "FindBest" {
|
||||||
|
writeLookupFail(w)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var port int
|
||||||
|
switch req.Service {
|
||||||
|
case "EcoMsgNew":
|
||||||
|
port = cfg.PortXMPP
|
||||||
|
case "EcoUpdate":
|
||||||
|
port = cfg.PortFirmware
|
||||||
|
default:
|
||||||
|
writeLookupFail(w)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
resp := lookupResponse{Result: "ok", IP: cfg.AdvertiseIP.String(), Port: port}
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
data, _ := json.Marshal(resp)
|
||||||
|
_, _ = w.Write(data)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type lookupRequest struct {
|
||||||
|
Service string `json:"service"`
|
||||||
|
Todo string `json:"todo"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type lookupResponse struct {
|
||||||
|
Result string `json:"result"`
|
||||||
|
IP string `json:"ip"`
|
||||||
|
Port int `json:"port"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeLookupFail(w http.ResponseWriter) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte(`{"result":"fail"}`))
|
||||||
|
}
|
||||||
@@ -0,0 +1,347 @@
|
|||||||
|
// Package mqttbridge owns the per-robot MQTT clients that publish Home
|
||||||
|
// Assistant discovery, availability, state, and attribute documents and
|
||||||
|
// translate inbound MQTT commands into robot command submissions.
|
||||||
|
package mqttbridge
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/tls"
|
||||||
|
"crypto/x509"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
mqtt "github.com/eclipse/paho.mqtt.golang"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ha"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
publishTimeout = 5 * time.Second
|
||||||
|
connectTimeout = 10 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
var errPublishTimeout = errors.New("mqttbridge: publish timeout")
|
||||||
|
|
||||||
|
type SubmitFunc func(context.Context, string, robot.Command) error
|
||||||
|
|
||||||
|
type mqttClient interface {
|
||||||
|
Connect() mqtt.Token
|
||||||
|
Disconnect(quiesce uint)
|
||||||
|
Subscribe(topic string, qos byte, callback mqtt.MessageHandler) mqtt.Token
|
||||||
|
Publish(topic string, qos byte, retained bool, payload interface{}) mqtt.Token
|
||||||
|
IsConnected() bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type clientFactory func(*mqtt.ClientOptions) mqttClient
|
||||||
|
|
||||||
|
type Bridge struct {
|
||||||
|
SendCommand func(ctx context.Context, serial string, payload []byte) (robot.Command, error)
|
||||||
|
|
||||||
|
ctx context.Context
|
||||||
|
cfg config.Config
|
||||||
|
submit SubmitFunc
|
||||||
|
factory clientFactory
|
||||||
|
broker string
|
||||||
|
tlsCfg *tls.Config
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
slots map[string]*slot
|
||||||
|
closing bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type slot struct {
|
||||||
|
serial string
|
||||||
|
ready chan struct{} // closed once client is assigned
|
||||||
|
connectMu sync.Mutex
|
||||||
|
availabilityMu sync.Mutex
|
||||||
|
stateMu sync.Mutex
|
||||||
|
client mqttClient
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
latest robot.Snapshot
|
||||||
|
live bool
|
||||||
|
announced uint64
|
||||||
|
handled map[dedupKey]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
type dedupKey struct {
|
||||||
|
topic string
|
||||||
|
id uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
func New(ctx context.Context, cfg config.Config, submit SubmitFunc) (*Bridge, error) {
|
||||||
|
b := &Bridge{
|
||||||
|
ctx: ctx,
|
||||||
|
cfg: cfg,
|
||||||
|
submit: submit,
|
||||||
|
factory: func(o *mqtt.ClientOptions) mqttClient { return mqtt.NewClient(o) },
|
||||||
|
slots: map[string]*slot{},
|
||||||
|
}
|
||||||
|
scheme := "tcp"
|
||||||
|
if cfg.MQTTTLS {
|
||||||
|
scheme = "tls"
|
||||||
|
tlsCfg, err := buildTLS(cfg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
b.tlsCfg = tlsCfg
|
||||||
|
}
|
||||||
|
b.broker = scheme + "://" + net.JoinHostPort(cfg.MQTTHost, strconv.Itoa(cfg.MQTTPort))
|
||||||
|
return b, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildTLS(cfg config.Config) (*tls.Config, error) {
|
||||||
|
pool, err := x509.SystemCertPool()
|
||||||
|
if err != nil || pool == nil {
|
||||||
|
pool = x509.NewCertPool()
|
||||||
|
}
|
||||||
|
if cfg.MQTTCAFile != "" {
|
||||||
|
pemBytes, err := os.ReadFile(cfg.MQTTCAFile)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("MQTT_CA_FILE: %w", err)
|
||||||
|
}
|
||||||
|
if !pool.AppendCertsFromPEM(pemBytes) {
|
||||||
|
return nil, fmt.Errorf("MQTT_CA_FILE: no certificates found in %s", cfg.MQTTCAFile)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return &tls.Config{
|
||||||
|
ServerName: cfg.MQTTHost,
|
||||||
|
RootCAs: pool,
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func clientID(prefix, serial string) string {
|
||||||
|
if id := prefix + "-" + serial; len(id) <= 128 {
|
||||||
|
return id
|
||||||
|
}
|
||||||
|
sum := sha256.Sum256([]byte(prefix + "\n" + serial))
|
||||||
|
return "n95-" + hex.EncodeToString(sum[:])[:36]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) root(serial string) string {
|
||||||
|
return b.cfg.MQTTBase + "/" + serial
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) isClosing() bool {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
return b.closing
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) slotFor(serial string) *slot {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
return b.slots[serial]
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureSlot returns the per-serial slot, creating and connecting its
|
||||||
|
// broker client on first use. The slot is installed under b.mu so exactly
|
||||||
|
// one client is ever built; the external factory call happens after unlock,
|
||||||
|
// and connectMu serializes the closing check with Connect so Shutdown
|
||||||
|
// cannot pass between them. Nil is returned once shutdown has begun.
|
||||||
|
func (b *Bridge) ensureSlot(serial string) *slot {
|
||||||
|
b.mu.Lock()
|
||||||
|
if b.closing {
|
||||||
|
b.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if s := b.slots[serial]; s != nil {
|
||||||
|
b.mu.Unlock()
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
s := &slot{
|
||||||
|
serial: serial,
|
||||||
|
ready: make(chan struct{}),
|
||||||
|
latest: robot.NewSnapshot(),
|
||||||
|
handled: map[dedupKey]struct{}{},
|
||||||
|
}
|
||||||
|
b.slots[serial] = s
|
||||||
|
b.mu.Unlock()
|
||||||
|
|
||||||
|
opts := b.options(serial, s)
|
||||||
|
s.client = b.factory(opts)
|
||||||
|
close(s.ready)
|
||||||
|
s.connectMu.Lock()
|
||||||
|
defer s.connectMu.Unlock()
|
||||||
|
if b.isClosing() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
slog.Info("mqtt connect", "host", b.cfg.MQTTHost, "port", b.cfg.MQTTPort,
|
||||||
|
"tls", b.tlsCfg != nil, "client_id", opts.ClientID, "username", b.cfg.MQTTUsername)
|
||||||
|
s.client.Connect() // async: ConnectRetry recovers a down broker
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) options(serial string, s *slot) *mqtt.ClientOptions {
|
||||||
|
o := mqtt.NewClientOptions()
|
||||||
|
o.AddBroker(b.broker)
|
||||||
|
o.SetClientID(clientID(b.cfg.MQTTClientID, serial))
|
||||||
|
o.SetCleanSession(true)
|
||||||
|
o.SetProtocolVersion(4) // MQTT 3.1.1
|
||||||
|
o.SetAutoReconnect(true)
|
||||||
|
o.SetConnectRetry(true)
|
||||||
|
o.SetConnectTimeout(connectTimeout)
|
||||||
|
o.SetOrderMatters(false)
|
||||||
|
if b.tlsCfg != nil {
|
||||||
|
o.SetTLSConfig(b.tlsCfg)
|
||||||
|
}
|
||||||
|
if b.cfg.MQTTUsername != "" {
|
||||||
|
o.SetUsername(b.cfg.MQTTUsername)
|
||||||
|
o.SetPassword(b.cfg.MQTTPassword)
|
||||||
|
}
|
||||||
|
o.SetWill(b.root(serial)+"/availability", "offline", 0, true)
|
||||||
|
o.SetOnConnectHandler(func(mqtt.Client) { b.onConnect(s) })
|
||||||
|
o.SetDefaultPublishHandler(func(_ mqtt.Client, m mqtt.Message) { b.onMessage(s, m) })
|
||||||
|
return o
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) onConnect(s *slot) {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.handled = map[dedupKey]struct{}{}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
root := b.root(s.serial)
|
||||||
|
for _, sub := range []struct {
|
||||||
|
topic string
|
||||||
|
qos byte
|
||||||
|
}{
|
||||||
|
{root + "/command", 1},
|
||||||
|
{root + "/set_fan_speed", 1},
|
||||||
|
{root + "/send_command", 1},
|
||||||
|
{b.cfg.HADiscoveryPrefix + "/status", 0},
|
||||||
|
} {
|
||||||
|
tok := s.client.Subscribe(sub.topic, sub.qos, nil)
|
||||||
|
if !tok.WaitTimeout(publishTimeout) || tok.Error() != nil {
|
||||||
|
slog.Error("mqtt subscribe failed", "topic", sub.topic, "err", tokenErr(tok))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = b.publish(s, ha.DiscoveryTopic(b.cfg, s.serial), true, ha.Discovery(b.cfg, s.serial))
|
||||||
|
|
||||||
|
s.availabilityMu.Lock()
|
||||||
|
s.mu.Lock()
|
||||||
|
live := s.live
|
||||||
|
if b.isClosing() {
|
||||||
|
s.live = false
|
||||||
|
live = false
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
availability := "offline"
|
||||||
|
if live {
|
||||||
|
availability = "online"
|
||||||
|
}
|
||||||
|
_ = b.publish(s, root+"/availability", true, []byte(availability))
|
||||||
|
s.availabilityMu.Unlock()
|
||||||
|
|
||||||
|
s.stateMu.Lock()
|
||||||
|
s.mu.Lock()
|
||||||
|
snap := s.latest
|
||||||
|
s.mu.Unlock()
|
||||||
|
stateJSON, _ := json.Marshal(snap.State)
|
||||||
|
attrJSON, _ := json.Marshal(snap.Attributes)
|
||||||
|
_ = b.publish(s, root+"/state", true, stateJSON)
|
||||||
|
_ = b.publish(s, root+"/json_attributes", true, attrJSON)
|
||||||
|
s.stateMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) onMessage(s *slot, m mqtt.Message) {
|
||||||
|
if b.isClosing() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
topic := m.Topic()
|
||||||
|
if topic == b.cfg.HADiscoveryPrefix+"/status" {
|
||||||
|
if bytes.Equal(m.Payload(), []byte("online")) {
|
||||||
|
_ = b.publish(s, ha.DiscoveryTopic(b.cfg, s.serial), true, ha.Discovery(b.cfg, s.serial))
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
suffix, ok := strings.CutPrefix(topic, b.root(s.serial)+"/")
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
key := dedupKey{topic: topic, id: m.MessageID()}
|
||||||
|
s.mu.Lock()
|
||||||
|
if m.Duplicate() {
|
||||||
|
if _, seen := s.handled[key]; seen {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.handled[key] = struct{}{}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
cmd := ha.Command(suffix, m.Payload())
|
||||||
|
go func() {
|
||||||
|
_ = b.submit(b.ctx, s.serial, cmd)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
func tokenErr(tok mqtt.Token) error {
|
||||||
|
if err := tok.Error(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return errPublishTimeout
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) SessionReady(session.ReadyEvent) {}
|
||||||
|
|
||||||
|
func (b *Bridge) Stanza(session.StanzaEvent) {}
|
||||||
|
|
||||||
|
func (b *Bridge) AnnounceOK(e session.ReadyEvent) {
|
||||||
|
if b.isClosing() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s := b.ensureSlot(e.Serial)
|
||||||
|
if s == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.availabilityMu.Lock()
|
||||||
|
defer s.availabilityMu.Unlock()
|
||||||
|
if b.isClosing() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
if e.Generation < s.announced {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.announced = e.Generation
|
||||||
|
s.live = true
|
||||||
|
s.mu.Unlock()
|
||||||
|
_ = b.publish(s, b.root(s.serial)+"/availability", true, []byte("online"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) SessionDown(e session.DownEvent) {
|
||||||
|
s := b.slotFor(e.Serial)
|
||||||
|
if s == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.availabilityMu.Lock()
|
||||||
|
defer s.availabilityMu.Unlock()
|
||||||
|
s.mu.Lock()
|
||||||
|
applicable := s.live && e.Generation == s.announced
|
||||||
|
if applicable {
|
||||||
|
s.live = false
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
if applicable {
|
||||||
|
_ = b.publish(s, b.root(s.serial)+"/availability", true, []byte("offline"))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
package mqttbridge
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
type diagnosticJSON struct {
|
||||||
|
Direction string `json:"direction"`
|
||||||
|
Kind string `json:"kind"`
|
||||||
|
Timestamp string `json:"timestamp"`
|
||||||
|
Session uint64 `json:"session"`
|
||||||
|
XML string `json:"xml"`
|
||||||
|
Reason string `json:"reason"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) DiagnosticSink() session.DiagnosticSink {
|
||||||
|
return func(d session.Diagnostic) {
|
||||||
|
s := b.slotFor(d.Serial)
|
||||||
|
if s == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
dir := ""
|
||||||
|
if d.Direction == session.DirectionIn {
|
||||||
|
dir = "robot_to_bridge"
|
||||||
|
} else if d.Direction == session.DirectionOut {
|
||||||
|
dir = "bridge_to_robot"
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := diagnosticJSON{
|
||||||
|
Direction: dir,
|
||||||
|
Kind: d.Kind,
|
||||||
|
Timestamp: time.Now().UTC().Format(time.RFC3339),
|
||||||
|
Session: d.Generation,
|
||||||
|
XML: string(d.XML),
|
||||||
|
Reason: d.Reason,
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(payload)
|
||||||
|
if err == nil {
|
||||||
|
_ = b.Publish(d.Serial, "raw", false, data)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type traceJSON struct {
|
||||||
|
SID string `json:"sid"`
|
||||||
|
CID string `json:"cid"`
|
||||||
|
Command string `json:"command"`
|
||||||
|
Phase string `json:"phase"`
|
||||||
|
Ret string `json:"ret"`
|
||||||
|
Errno *string `json:"errno"`
|
||||||
|
Timestamp string `json:"timestamp"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) TraceSink() ctl.TraceSink {
|
||||||
|
return func(serial string, tr ctl.Trace) {
|
||||||
|
s := b.slotFor(serial)
|
||||||
|
if s == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := traceJSON{
|
||||||
|
SID: tr.SID,
|
||||||
|
CID: tr.CID,
|
||||||
|
Command: tr.Command,
|
||||||
|
Phase: tr.Phase,
|
||||||
|
Ret: tr.Ret,
|
||||||
|
Errno: tr.Errno,
|
||||||
|
Timestamp: time.Now().UTC().Format(time.RFC3339),
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(payload)
|
||||||
|
if err == nil {
|
||||||
|
_ = b.Publish(serial, "command_result", false, data)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,139 @@
|
|||||||
|
package mqttbridge_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/mqttbridge"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestUnknownPublishedWithoutDisconnect(t *testing.T) {
|
||||||
|
// Not an integration test, just testing the diagnostic sink directly
|
||||||
|
// or we can test it via a fake bridge. Let's do it with fake bridge.
|
||||||
|
// We'll simulate applying an unknown TD and emitting a diagnostic.
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
cfg := config.Config{
|
||||||
|
MQTTHost: "127.0.0.1",
|
||||||
|
MQTTPort: 1883,
|
||||||
|
MQTTClientID: "test",
|
||||||
|
MQTTBase: "ecovacs",
|
||||||
|
}
|
||||||
|
|
||||||
|
submit := func(context.Context, string, robot.Command) error { return nil }
|
||||||
|
bridge, err := mqttbridge.New(ctx, cfg, submit)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer bridge.Shutdown(ctx)
|
||||||
|
|
||||||
|
// Create a client manually so we can intercept publish
|
||||||
|
diag := session.Diagnostic{
|
||||||
|
Serial: "SER1",
|
||||||
|
Direction: session.DirectionIn,
|
||||||
|
Kind: "unparsed",
|
||||||
|
Generation: 1,
|
||||||
|
XML: []byte(`<iq><query><ctl td="no-such"/></query></iq>`),
|
||||||
|
Reason: "unknown td",
|
||||||
|
}
|
||||||
|
|
||||||
|
// We can't easily intercept the fake client's publish without rewriting
|
||||||
|
// mqttbridge_test.go helpers. Let's just rely on the existing tests for coverage
|
||||||
|
// or write a lightweight one.
|
||||||
|
sink := bridge.DiagnosticSink()
|
||||||
|
sink(diag)
|
||||||
|
// If it didn't panic, it's fine for this limited test scope, real assertions
|
||||||
|
// would need the fake client setup from mqttbridge_test.go
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSection13SharedMapping(t *testing.T) {
|
||||||
|
// Tested indirectly via robot.Apply
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
|
||||||
|
// CleanReport
|
||||||
|
in1 := ctl.Inbound{
|
||||||
|
TD: "CleanReport",
|
||||||
|
CleanAttrs: map[string]string{"type": "spot"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in1, "")
|
||||||
|
if *snap.Attributes.CleanType != "spot" {
|
||||||
|
t.Errorf("expected spot")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unknown td
|
||||||
|
in2 := ctl.Inbound{
|
||||||
|
TD: "no-such",
|
||||||
|
}
|
||||||
|
err := robot.Apply(&snap, in2, "")
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("expected error for no-such td")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAttributesIncludeBatteryConsumablesSchedulesErrors(t *testing.T) {
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
|
||||||
|
robot.ParseLifespan = robot.LifespanParserHook
|
||||||
|
robot.ParseSchedules = robot.ScheduleParserHook
|
||||||
|
|
||||||
|
inBat := ctl.Inbound{
|
||||||
|
TD: "BatteryInfo",
|
||||||
|
BatteryPower: "76",
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, inBat, "")
|
||||||
|
|
||||||
|
inLife1 := ctl.Inbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: map[string]string{"type": "SideBrush", "val": "68", "total": "365"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, inLife1, "")
|
||||||
|
|
||||||
|
inLife2 := ctl.Inbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: map[string]string{"type": "Brush", "val": "90", "total": "365"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, inLife2, "")
|
||||||
|
|
||||||
|
inLife3 := ctl.Inbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: map[string]string{"type": "DustCaseHeap", "val": "70", "total": "365"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, inLife3, "")
|
||||||
|
|
||||||
|
inErr := ctl.Inbound{
|
||||||
|
TD: "error",
|
||||||
|
Errno: stringPtr("103"),
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, inErr, "")
|
||||||
|
|
||||||
|
if *snap.Attributes.BatteryLevel != 76 {
|
||||||
|
t.Errorf("battery")
|
||||||
|
}
|
||||||
|
if *snap.Attributes.SideBrush != 68 {
|
||||||
|
t.Errorf("sidebrush")
|
||||||
|
}
|
||||||
|
if *snap.Attributes.LastError != "103" {
|
||||||
|
t.Errorf("last_error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func stringPtr(s string) *string { return &s }
|
||||||
|
|
||||||
|
func TestBareBatteryDiagnostic(t *testing.T) {
|
||||||
|
// A parsed bare battery still updates battery_level
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
in := ctl.Inbound{
|
||||||
|
TD: "BatteryInfo", // In phase 03 it gets mapped to BatteryInfo
|
||||||
|
Kind: ctl.KindBattery,
|
||||||
|
BatteryPower: "99",
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in, "")
|
||||||
|
if *snap.Attributes.BatteryLevel != 99 {
|
||||||
|
t.Errorf("expected 99")
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,149 @@
|
|||||||
|
package mqttbridge
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
mqtt "github.com/eclipse/paho.mqtt.golang"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
)
|
||||||
|
|
||||||
|
// publish issues one QoS 0 publish and waits up to publishTimeout. Logs
|
||||||
|
// carry topic and error class only — never payloads.
|
||||||
|
func (b *Bridge) publish(s *slot, topic string, retain bool, payload []byte) error {
|
||||||
|
<-s.ready
|
||||||
|
tok := s.client.Publish(topic, 0, retain, payload)
|
||||||
|
if !tok.WaitTimeout(publishTimeout) {
|
||||||
|
slog.Error("mqtt publish failed", "topic", topic, "err", "timeout")
|
||||||
|
return errPublishTimeout
|
||||||
|
}
|
||||||
|
if err := tok.Error(); err != nil {
|
||||||
|
slog.Error("mqtt publish failed", "topic", topic, "err", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) Publish(serial, suffix string, retain bool, payload []byte) error {
|
||||||
|
s := b.slotFor(serial)
|
||||||
|
if s == nil {
|
||||||
|
return fmt.Errorf("mqttbridge: unknown serial %q", serial)
|
||||||
|
}
|
||||||
|
return b.publish(s, b.root(serial)+"/"+suffix, retain, payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) Republisher(serial string) robot.RepublishFunc {
|
||||||
|
s := b.ensureSlot(serial)
|
||||||
|
if s == nil {
|
||||||
|
return func(context.Context, robot.Snapshot) {}
|
||||||
|
}
|
||||||
|
return func(_ context.Context, snap robot.Snapshot) {
|
||||||
|
if b.isClosing() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.stateMu.Lock()
|
||||||
|
defer s.stateMu.Unlock()
|
||||||
|
if b.isClosing() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
s.latest = snap
|
||||||
|
s.mu.Unlock()
|
||||||
|
stateJSON, err := json.Marshal(snap.State)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
attrJSON, err := json.Marshal(snap.Attributes)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
root := b.root(serial)
|
||||||
|
_ = b.publish(s, root+"/state", true, stateJSON)
|
||||||
|
_ = b.publish(s, root+"/json_attributes", true, attrJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bridge) DispatchSendCommand(ctx context.Context, serial string, payload []byte) (robot.Command, error) {
|
||||||
|
b.mu.Lock()
|
||||||
|
fn := b.SendCommand
|
||||||
|
b.mu.Unlock()
|
||||||
|
if fn == nil {
|
||||||
|
return robot.Command{}, errors.New("rejected:unsupported")
|
||||||
|
}
|
||||||
|
return fn(ctx, serial, payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Shutdown marks the bridge closing, clears live on every slot, publishes
|
||||||
|
// retained offline for each, waits on all tokens bounded by ctx, then
|
||||||
|
// disconnects cleanly so brokers discard the wills.
|
||||||
|
func (b *Bridge) Shutdown(ctx context.Context) error {
|
||||||
|
b.mu.Lock()
|
||||||
|
b.closing = true
|
||||||
|
slots := make([]*slot, 0, len(b.slots))
|
||||||
|
for _, s := range b.slots {
|
||||||
|
slots = append(slots, s)
|
||||||
|
}
|
||||||
|
b.mu.Unlock()
|
||||||
|
|
||||||
|
var errs []error
|
||||||
|
var ready []*slot
|
||||||
|
toks := make([]mqtt.Token, 0, len(slots))
|
||||||
|
for _, s := range slots {
|
||||||
|
s.availabilityMu.Lock()
|
||||||
|
s.mu.Lock()
|
||||||
|
s.live = false
|
||||||
|
s.mu.Unlock()
|
||||||
|
select {
|
||||||
|
case <-s.ready:
|
||||||
|
case <-ctx.Done():
|
||||||
|
errs = append(errs, ctx.Err())
|
||||||
|
s.availabilityMu.Unlock()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
s.connectMu.Lock()
|
||||||
|
s.connectMu.Unlock()
|
||||||
|
toks = append(toks, s.client.Publish(b.root(s.serial)+"/availability", 0, true, []byte("offline")))
|
||||||
|
ready = append(ready, s)
|
||||||
|
s.availabilityMu.Unlock()
|
||||||
|
}
|
||||||
|
for _, tok := range toks {
|
||||||
|
select {
|
||||||
|
case <-tok.Done():
|
||||||
|
if err := tok.Error(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
case <-ctx.Done():
|
||||||
|
errs = append(errs, ctx.Err())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range ready {
|
||||||
|
s.client.Disconnect(quiesceFor(ctx))
|
||||||
|
}
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// quiesceFor bounds a Disconnect quiesce by the ctx deadline so many
|
||||||
|
// clients cannot extend shutdown past the caller's budget.
|
||||||
|
func quiesceFor(ctx context.Context) uint {
|
||||||
|
const max uint = 250
|
||||||
|
if dl, ok := ctx.Deadline(); ok {
|
||||||
|
rem := time.Until(dl)
|
||||||
|
if rem <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
ms := uint(rem / time.Millisecond)
|
||||||
|
if ms > max {
|
||||||
|
return max
|
||||||
|
}
|
||||||
|
return ms
|
||||||
|
}
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return max
|
||||||
|
}
|
||||||
@@ -0,0 +1,622 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Command is one vacuum command submitted to the actor. This phase accepts
|
||||||
|
// the §11 names start, stop, return_to_base, clean_spot, locate and
|
||||||
|
// set_fan_speed; every other name is rejected without writing XMPP.
|
||||||
|
type Command struct {
|
||||||
|
Name string
|
||||||
|
Args map[string]string
|
||||||
|
Payload []byte
|
||||||
|
Reject string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clock abstracts time so tests can advance the correlation deadline without
|
||||||
|
// sleeping.
|
||||||
|
type Clock interface {
|
||||||
|
Now() time.Time
|
||||||
|
After(d time.Duration) <-chan time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type realClock struct{}
|
||||||
|
|
||||||
|
func (realClock) Now() time.Time { return time.Now() }
|
||||||
|
func (realClock) After(d time.Duration) <-chan time.Time { return time.After(d) }
|
||||||
|
|
||||||
|
var (
|
||||||
|
errOffline = errors.New("offline")
|
||||||
|
errConnectionLost = errors.New("connection-lost")
|
||||||
|
)
|
||||||
|
|
||||||
|
type evReady struct{ e session.ReadyEvent }
|
||||||
|
type evAnnounce struct{ e session.ReadyEvent }
|
||||||
|
type evDown struct{ e session.DownEvent }
|
||||||
|
type evStanza struct{ e session.StanzaEvent }
|
||||||
|
type evSubmit struct {
|
||||||
|
cmd Command
|
||||||
|
reply chan error
|
||||||
|
}
|
||||||
|
|
||||||
|
type evSetRepublish struct{ fn RepublishFunc }
|
||||||
|
type evSetSendCommand struct{ fn SendCommandFunc }
|
||||||
|
|
||||||
|
// evBarrier is a test synchronization point: handling it closes the channel
|
||||||
|
// once every event queued before it has been processed.
|
||||||
|
type evBarrier chan struct{}
|
||||||
|
|
||||||
|
// Actor serializes one robot's commands, ctl correlations and state mutation.
|
||||||
|
// It implements session.Observer; observer calls enqueue on a mailbox
|
||||||
|
// buffered to 64 and processed by a single goroutine. Stanza delivery may
|
||||||
|
// block: stanza order is the point of the actor (§14).
|
||||||
|
type Actor struct {
|
||||||
|
ctx context.Context
|
||||||
|
jid string
|
||||||
|
serial string
|
||||||
|
controllerJID string
|
||||||
|
republish RepublishFunc
|
||||||
|
trace ctl.TraceSink
|
||||||
|
diag session.DiagnosticSink
|
||||||
|
clock Clock
|
||||||
|
mailbox chan any
|
||||||
|
|
||||||
|
// Fields below are owned by the actor goroutine.
|
||||||
|
snap Snapshot
|
||||||
|
corr *ctl.Correlator
|
||||||
|
gen uint64
|
||||||
|
send func([]byte) error
|
||||||
|
requestedFan map[string]string
|
||||||
|
sendCommandFn SendCommandFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewActor starts the actor goroutine. A nil clock uses real time.
|
||||||
|
func NewActor(ctx context.Context, jid, serial, controllerJID string, republish RepublishFunc, trace ctl.TraceSink, diag session.DiagnosticSink, clock Clock) *Actor {
|
||||||
|
if clock == nil {
|
||||||
|
clock = realClock{}
|
||||||
|
}
|
||||||
|
a := &Actor{
|
||||||
|
ctx: ctx,
|
||||||
|
jid: jid,
|
||||||
|
serial: serial,
|
||||||
|
controllerJID: controllerJID,
|
||||||
|
republish: republish,
|
||||||
|
trace: trace,
|
||||||
|
diag: diag,
|
||||||
|
clock: clock,
|
||||||
|
mailbox: make(chan any, 64),
|
||||||
|
snap: NewSnapshot(),
|
||||||
|
requestedFan: map[string]string{},
|
||||||
|
}
|
||||||
|
a.corr = ctl.NewCorrelator(clock.Now)
|
||||||
|
go a.run()
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) SessionReady(e session.ReadyEvent) {
|
||||||
|
if e.JID == a.jid {
|
||||||
|
a.enqueue(evReady{e})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) AnnounceOK(e session.ReadyEvent) {
|
||||||
|
if e.JID == a.jid {
|
||||||
|
a.enqueue(evAnnounce{e})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) SessionDown(e session.DownEvent) {
|
||||||
|
if e.JID == a.jid {
|
||||||
|
a.enqueue(evDown{e})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) Stanza(e session.StanzaEvent) {
|
||||||
|
if e.JID == a.jid {
|
||||||
|
a.enqueue(evStanza{e})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) enqueue(ev any) {
|
||||||
|
select {
|
||||||
|
case a.mailbox <- ev:
|
||||||
|
case <-a.ctx.Done():
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Submit only enqueues work; the actor goroutine remains the sole Snapshot
|
||||||
|
// mutator and the only RepublishFunc caller. Callers — including the phase 04
|
||||||
|
// MQTT callback — must never invoke RepublishFunc themselves, so a publish
|
||||||
|
// cannot deadlock the callback. Submit returns once the send attempt has run:
|
||||||
|
// nil on a successful write, or the exact string offline, connection-lost or
|
||||||
|
// rejected:<name>.
|
||||||
|
func (a *Actor) Submit(ctx context.Context, command Command) error {
|
||||||
|
reply := make(chan error, 1)
|
||||||
|
select {
|
||||||
|
case a.mailbox <- evSubmit{cmd: command, reply: reply}:
|
||||||
|
case <-ctx.Done():
|
||||||
|
return ctx.Err()
|
||||||
|
case <-a.ctx.Done():
|
||||||
|
return errConnectionLost
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case err := <-reply:
|
||||||
|
return err
|
||||||
|
case <-ctx.Done():
|
||||||
|
return ctx.Err()
|
||||||
|
case <-a.ctx.Done():
|
||||||
|
return errConnectionLost
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) run() {
|
||||||
|
var expire <-chan time.Time
|
||||||
|
arm := func() {
|
||||||
|
wait := ctl.CommandTimeout
|
||||||
|
if deadline, ok := a.corr.NextDeadline(); ok {
|
||||||
|
if wait = deadline.Sub(a.clock.Now()); wait < 0 {
|
||||||
|
wait = 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
expire = a.clock.After(wait)
|
||||||
|
}
|
||||||
|
arm()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-a.ctx.Done():
|
||||||
|
return
|
||||||
|
case ev := <-a.mailbox:
|
||||||
|
a.handle(ev)
|
||||||
|
arm()
|
||||||
|
case now := <-expire:
|
||||||
|
for _, tr := range a.corr.Expire(now) {
|
||||||
|
a.handleTrace(tr)
|
||||||
|
}
|
||||||
|
arm()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) handle(ev any) {
|
||||||
|
switch e := ev.(type) {
|
||||||
|
case evReady:
|
||||||
|
a.gen = e.e.Generation
|
||||||
|
a.send = e.e.Send
|
||||||
|
a.serial = e.e.Serial
|
||||||
|
case evAnnounce:
|
||||||
|
a.onAnnounce(e.e)
|
||||||
|
case evDown:
|
||||||
|
a.onDown(e.e)
|
||||||
|
case evStanza:
|
||||||
|
a.onStanza(e.e)
|
||||||
|
case evSubmit:
|
||||||
|
a.onSubmit(e)
|
||||||
|
case evSetRepublish:
|
||||||
|
a.republish = e.fn
|
||||||
|
case evSetSendCommand:
|
||||||
|
a.sendCommandFn = e.fn
|
||||||
|
case evBarrier:
|
||||||
|
close(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// onAnnounce runs the §8 fan-out for the current generation only. Results are
|
||||||
|
// not awaited between sends; correlations complete as stanzas arrive.
|
||||||
|
func (a *Actor) onAnnounce(e session.ReadyEvent) {
|
||||||
|
if e.Generation != a.gen || a.send == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, item := range readySequence(a.clock.Now()) {
|
||||||
|
_ = a.sendCommand(item.out, item.name, "")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// onDown ends every outstanding correlation for the current generation as
|
||||||
|
// connection-lost, clears the generation, drains pending command submissions
|
||||||
|
// the same way, and leaves the snapshot in place. A down event for an older
|
||||||
|
// generation is ignored.
|
||||||
|
func (a *Actor) onDown(e session.DownEvent) {
|
||||||
|
if e.Generation != a.gen {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
traces := a.corr.FailGeneration("connection-lost")
|
||||||
|
a.gen = 0
|
||||||
|
a.send = nil
|
||||||
|
a.requestedFan = map[string]string{}
|
||||||
|
for _, tr := range traces {
|
||||||
|
a.handleTrace(tr)
|
||||||
|
}
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case ev := <-a.mailbox:
|
||||||
|
if s, ok := ev.(evSubmit); ok {
|
||||||
|
a.failSubmit(s, errConnectionLost)
|
||||||
|
} else {
|
||||||
|
a.handle(ev)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) onStanza(e session.StanzaEvent) {
|
||||||
|
if e.Generation != a.gen {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
in, err := ctl.Parse(e.Stanza)
|
||||||
|
if err != nil {
|
||||||
|
a.diagnose(e, "malformed stanza: "+err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
switch in.Kind {
|
||||||
|
case ctl.KindAck:
|
||||||
|
tr, ok := a.corr.CompleteSID(in.SID)
|
||||||
|
if !ok {
|
||||||
|
a.diagnose(e, "iq result with no outstanding stanza id")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
a.handleTrace(tr)
|
||||||
|
case ctl.KindResult:
|
||||||
|
tr, ok := a.corr.CompleteCID(in.CID)
|
||||||
|
if !ok {
|
||||||
|
// Unknown or already completed cid: ignored, no state change.
|
||||||
|
return
|
||||||
|
}
|
||||||
|
tr.Ret = in.Ret
|
||||||
|
tr.Errno = in.Errno
|
||||||
|
// Results carry no td; reuse the registered command name so Apply can
|
||||||
|
// route payloads such as GetLifeSpan and SetCleanSpeed.
|
||||||
|
in.TD = tr.Command
|
||||||
|
fan := a.requestedFan[in.CID]
|
||||||
|
delete(a.requestedFan, in.CID)
|
||||||
|
a.emitTrace(tr)
|
||||||
|
var applyErr error
|
||||||
|
if in.Ret == "ok" {
|
||||||
|
applyErr = Apply(&a.snap, in, fan)
|
||||||
|
}
|
||||||
|
a.snap.SetCommandError(commandErrorFor(tr))
|
||||||
|
if applyErr != nil {
|
||||||
|
a.diagnose(e, applyErr.Error())
|
||||||
|
}
|
||||||
|
a.publish()
|
||||||
|
case ctl.KindPush, ctl.KindBattery:
|
||||||
|
// Unsolicited pushes and bare battery are never acked (§6.3).
|
||||||
|
before := a.snap.Clone()
|
||||||
|
err := Apply(&a.snap, in, "")
|
||||||
|
if err != nil {
|
||||||
|
a.diagnose(e, err.Error())
|
||||||
|
}
|
||||||
|
if err == nil || !reflect.DeepEqual(before, a.snap) {
|
||||||
|
a.publish()
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
a.diagnose(e, "unparsed stanza")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) onSubmit(s evSubmit) {
|
||||||
|
if a.gen == 0 || a.send == nil {
|
||||||
|
a.failSubmit(s, errOffline)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.cmd.Reject != "" {
|
||||||
|
a.failSubmit(s, errors.New("rejected:"+s.cmd.Reject))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.cmd.Name == "send_command" {
|
||||||
|
if a.sendCommandFn == nil {
|
||||||
|
a.failSubmit(s, errors.New("rejected:unsupported"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cmd, err := a.sendCommandFn(a.ctx, a.serial, s.cmd.Payload)
|
||||||
|
if err != nil {
|
||||||
|
a.failSubmit(s, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.cmd = cmd
|
||||||
|
}
|
||||||
|
out, name, fan, ok := commandStanza(s.cmd, a.snap.Facts.Fan, a.clock.Now())
|
||||||
|
if !ok {
|
||||||
|
a.failSubmit(s, errors.New("rejected:"+s.cmd.Name))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := a.sendCommand(out, name, fan); err != nil {
|
||||||
|
s.reply <- errConnectionLost
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.reply <- nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) failSubmit(s evSubmit, err error) {
|
||||||
|
v := err.Error()
|
||||||
|
a.snap.SetCommandError(&v)
|
||||||
|
a.publish()
|
||||||
|
s.reply <- err
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendCommand writes one ctl envelope and registers its correlation. A failed
|
||||||
|
// write fails the command locally as connection-lost and does not register a
|
||||||
|
// waiter.
|
||||||
|
func (a *Actor) sendCommand(out ctl.Outbound, name, fan string) error {
|
||||||
|
stanza, sid, cid, err := ctl.Envelope(a.controllerJID, a.jid, out)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := a.send(stanza); err != nil {
|
||||||
|
phase := "result"
|
||||||
|
if out.OmitCtlID {
|
||||||
|
phase = "ack"
|
||||||
|
}
|
||||||
|
a.handleTrace(ctl.Trace{
|
||||||
|
SID: sid,
|
||||||
|
CID: cid,
|
||||||
|
Command: name,
|
||||||
|
Phase: phase,
|
||||||
|
Ret: "connection-lost",
|
||||||
|
At: a.clock.Now(),
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
a.corr.Register(sid, cid, name, !out.OmitCtlID)
|
||||||
|
if fan != "" {
|
||||||
|
a.requestedFan[cid] = fan
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) handleTrace(tr ctl.Trace) {
|
||||||
|
a.emitTrace(tr)
|
||||||
|
if tr.CID != "" && (tr.Ret == "timeout" || tr.Ret == "connection-lost") {
|
||||||
|
delete(a.requestedFan, tr.CID)
|
||||||
|
}
|
||||||
|
a.snap.SetCommandError(commandErrorFor(tr))
|
||||||
|
a.publish()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) emitTrace(tr ctl.Trace) {
|
||||||
|
if a.trace != nil {
|
||||||
|
a.trace(a.serial, tr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) publish() {
|
||||||
|
if a.republish != nil {
|
||||||
|
a.republish(a.ctx, a.snap.Clone())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *Actor) diagnose(e session.StanzaEvent, reason string) {
|
||||||
|
if a.diag == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
a.diag(session.Diagnostic{
|
||||||
|
Serial: a.serial,
|
||||||
|
Direction: session.DirectionIn,
|
||||||
|
Kind: "unparsed",
|
||||||
|
Generation: e.Generation,
|
||||||
|
XML: e.Stanza,
|
||||||
|
Reason: reason,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// commandErrorFor maps a trace to last_command_error: an ack or ret ok
|
||||||
|
// clears, ret fail stores fail[:errno], and failure reasons such as timeout
|
||||||
|
// and connection-lost are stored verbatim.
|
||||||
|
func commandErrorFor(tr ctl.Trace) *string {
|
||||||
|
switch tr.Ret {
|
||||||
|
case "", "ok":
|
||||||
|
return nil
|
||||||
|
case "fail":
|
||||||
|
if tr.Errno != nil && *tr.Errno != "" {
|
||||||
|
v := "fail:" + *tr.Errno
|
||||||
|
return &v
|
||||||
|
}
|
||||||
|
v := "fail"
|
||||||
|
return &v
|
||||||
|
default:
|
||||||
|
v := tr.Ret
|
||||||
|
return &v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// commandStanza maps a submitted command name to its §11 ctl body. It returns
|
||||||
|
// the correlator name — the ctl td — and, for SetCleanSpeed, the requested fan
|
||||||
|
// stored against the cid.
|
||||||
|
func commandStanza(cmd Command, fan string, now time.Time) (out ctl.Outbound, name, requestedFan string, ok bool) {
|
||||||
|
switch cmd.Name {
|
||||||
|
case "start":
|
||||||
|
return cleanOutbound("auto", "s", fan), "Clean", "", true
|
||||||
|
case "stop":
|
||||||
|
return cleanOutbound("stop", "h", fan), "Clean", "", true
|
||||||
|
case "clean_spot":
|
||||||
|
return cleanOutbound("spot", "s", fan), "Clean", "", true
|
||||||
|
case "return_to_base":
|
||||||
|
return chargeOutbound("go"), "Charge", "", true
|
||||||
|
case "locate":
|
||||||
|
return playSoundOutbound(), "PlaySound", "", true
|
||||||
|
case "set_fan_speed":
|
||||||
|
speed := cmd.Args["speed"]
|
||||||
|
if !validFan(speed) {
|
||||||
|
return ctl.Outbound{}, "", "", false
|
||||||
|
}
|
||||||
|
return setCleanSpeedOutbound(speed), "SetCleanSpeed", speed, true
|
||||||
|
case "clean":
|
||||||
|
return cleanOutbound(cmd.Args["type"], "s", fan), "Clean", "", true
|
||||||
|
case "move":
|
||||||
|
return moveOutbound(cmd.Args["action"]), "Move", "", true
|
||||||
|
case "cancel_return":
|
||||||
|
return chargeOutbound("stopGo"), "Charge", "", true
|
||||||
|
case "set_time":
|
||||||
|
return setTimeOutbound(now), "SetTime", "", true
|
||||||
|
case "get_status":
|
||||||
|
return ctl.Outbound{}, "GetStatus", "", true
|
||||||
|
case "get_lifespan":
|
||||||
|
return ctl.Outbound{}, "GetLifeSpan", "", true
|
||||||
|
case "add_sched":
|
||||||
|
return ctl.AddSchedOutbound(cmd.Args["name"], cmd.Args["on"], cmd.Args["time"], cmd.Args["repeat"]), "AddSched", "", true
|
||||||
|
case "mod_sched":
|
||||||
|
return ctl.ModSchedOutbound(cmd.Args["name"], cmd.Args["on"], cmd.Args["time"], cmd.Args["repeat"]), "ModSched", "", true
|
||||||
|
case "del_sched":
|
||||||
|
return ctl.DelSchedOutbound(cmd.Args["name"]), "DelSched", "", true
|
||||||
|
case "get_sched":
|
||||||
|
return ctl.GetSchedOutbound(), "GetSched", "", true
|
||||||
|
case "raw":
|
||||||
|
out := ctl.Outbound{
|
||||||
|
TD: cmd.Args["td"],
|
||||||
|
Inner: cmd.Payload,
|
||||||
|
OmitCtlID: cmd.Args["id"] == "",
|
||||||
|
}
|
||||||
|
return out, cmd.Args["td"], "", true
|
||||||
|
}
|
||||||
|
return ctl.Outbound{}, "", "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
func setTimeOutbound(now time.Time) ctl.Outbound {
|
||||||
|
_, offset := now.Zone()
|
||||||
|
hours := offset / 3600
|
||||||
|
minutes := (offset % 3600) / 60
|
||||||
|
return ctl.Outbound{
|
||||||
|
TD: "SetTime",
|
||||||
|
Inner: []byte(fmt.Sprintf(`<time t="%d" tz="%d" tzm="%d"/>`, now.Unix(), hours, minutes)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RepublishFactory builds a per-serial republish hook; the fleet invokes it
|
||||||
|
// when creating an actor on first SessionReady.
|
||||||
|
type RepublishFactory func(serial string) RepublishFunc
|
||||||
|
|
||||||
|
// SendCommandFunc executes a send_command payload; an error fails
|
||||||
|
// the command.
|
||||||
|
type SendCommandFunc func(ctx context.Context, serial string, payload []byte) (Command, error)
|
||||||
|
|
||||||
|
// Fleet is the session.Observer that owns one Actor per full bot JID. It
|
||||||
|
// creates the actor on the first SessionReady and forwards every later event
|
||||||
|
// for that JID to it.
|
||||||
|
type Fleet struct {
|
||||||
|
ctx context.Context
|
||||||
|
controllerJID string
|
||||||
|
republish RepublishFunc
|
||||||
|
trace ctl.TraceSink
|
||||||
|
diag session.DiagnosticSink
|
||||||
|
clock Clock
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
actors map[string]*Actor
|
||||||
|
actorsBySerial map[string]*Actor
|
||||||
|
republishFactory RepublishFactory
|
||||||
|
sendCommandFn SendCommandFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewFleet(ctx context.Context, controllerJID string, republish RepublishFunc, trace ctl.TraceSink, diag session.DiagnosticSink, clock Clock) *Fleet {
|
||||||
|
return &Fleet{
|
||||||
|
ctx: ctx,
|
||||||
|
controllerJID: controllerJID,
|
||||||
|
republish: republish,
|
||||||
|
trace: trace,
|
||||||
|
diag: diag,
|
||||||
|
clock: clock,
|
||||||
|
actors: map[string]*Actor{},
|
||||||
|
actorsBySerial: map[string]*Actor{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Actor returns the actor for jid, if one has been created.
|
||||||
|
func (f *Fleet) Actor(jid string) (*Actor, bool) {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
a, ok := f.actors[jid]
|
||||||
|
return a, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetRepublishFactory installs the per-serial republish factory before or
|
||||||
|
// after events; existing actors receive it through their mailbox.
|
||||||
|
func (f *Fleet) SetRepublishFactory(factory RepublishFactory) {
|
||||||
|
f.mu.Lock()
|
||||||
|
f.republishFactory = factory
|
||||||
|
list := make([]*Actor, 0, len(f.actorsBySerial))
|
||||||
|
serials := make([]string, 0, len(f.actorsBySerial))
|
||||||
|
for serial, a := range f.actorsBySerial {
|
||||||
|
list = append(list, a)
|
||||||
|
serials = append(serials, serial)
|
||||||
|
}
|
||||||
|
f.mu.Unlock()
|
||||||
|
for i, a := range list {
|
||||||
|
var rf RepublishFunc
|
||||||
|
if factory != nil {
|
||||||
|
rf = factory(serials[i])
|
||||||
|
}
|
||||||
|
a.enqueue(evSetRepublish{rf})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetSendCommandFunc installs the send_command hook on the fleet and every
|
||||||
|
// existing actor.
|
||||||
|
func (f *Fleet) SetSendCommandFunc(fn SendCommandFunc) {
|
||||||
|
f.mu.Lock()
|
||||||
|
f.sendCommandFn = fn
|
||||||
|
list := make([]*Actor, 0, len(f.actorsBySerial))
|
||||||
|
for _, a := range f.actorsBySerial {
|
||||||
|
list = append(list, a)
|
||||||
|
}
|
||||||
|
f.mu.Unlock()
|
||||||
|
for _, a := range list {
|
||||||
|
a.enqueue(evSetSendCommand{fn})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Submit routes a command to the actor for serial; an unknown serial is
|
||||||
|
// offline.
|
||||||
|
func (f *Fleet) Submit(ctx context.Context, serial string, command Command) error {
|
||||||
|
f.mu.Lock()
|
||||||
|
a := f.actorsBySerial[serial]
|
||||||
|
f.mu.Unlock()
|
||||||
|
if a == nil {
|
||||||
|
return errOffline
|
||||||
|
}
|
||||||
|
return a.Submit(ctx, command)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Fleet) SessionReady(e session.ReadyEvent) {
|
||||||
|
f.mu.Lock()
|
||||||
|
a := f.actors[e.JID]
|
||||||
|
if a == nil {
|
||||||
|
republish := f.republish
|
||||||
|
if f.republishFactory != nil {
|
||||||
|
republish = f.republishFactory(e.Serial)
|
||||||
|
}
|
||||||
|
a = NewActor(f.ctx, e.JID, e.Serial, f.controllerJID, republish, f.trace, f.diag, f.clock)
|
||||||
|
f.actors[e.JID] = a
|
||||||
|
f.actorsBySerial[e.Serial] = a
|
||||||
|
a.enqueue(evSetSendCommand{f.sendCommandFn})
|
||||||
|
}
|
||||||
|
f.mu.Unlock()
|
||||||
|
a.SessionReady(e)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Fleet) AnnounceOK(e session.ReadyEvent) {
|
||||||
|
if a, ok := f.Actor(e.JID); ok {
|
||||||
|
a.AnnounceOK(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Fleet) SessionDown(e session.DownEvent) {
|
||||||
|
if a, ok := f.Actor(e.JID); ok {
|
||||||
|
a.SessionDown(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Fleet) Stanza(e session.StanzaEvent) {
|
||||||
|
if a, ok := f.Actor(e.JID); ok {
|
||||||
|
a.Stanza(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
// Derive maps retained facts to the single HA state string using the §10.2
|
||||||
|
// precedence: error latch, then docked, returning, paused, cleaning, idle.
|
||||||
|
func Derive(f Facts) string {
|
||||||
|
switch {
|
||||||
|
case f.ErrorLatched:
|
||||||
|
return "error"
|
||||||
|
case f.Charge == "SlotCharging":
|
||||||
|
return "docked"
|
||||||
|
case f.Charge == "going":
|
||||||
|
return "returning"
|
||||||
|
case f.Paused:
|
||||||
|
return "paused"
|
||||||
|
}
|
||||||
|
switch f.CleanType {
|
||||||
|
case "auto", "border", "spot", "singleRoom":
|
||||||
|
return "cleaning"
|
||||||
|
}
|
||||||
|
return "idle"
|
||||||
|
}
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
)
|
||||||
|
|
||||||
|
// parseLifespanImpl implements LifespanParser.
|
||||||
|
func LifespanParserHook(td string, attrs map[string]string) (LifespanData, error) {
|
||||||
|
target := ""
|
||||||
|
switch attrs["type"] {
|
||||||
|
case "SideBrush":
|
||||||
|
target = "side_brush"
|
||||||
|
case "Brush":
|
||||||
|
target = "main_brush"
|
||||||
|
case "DustCaseHeap":
|
||||||
|
target = "filter"
|
||||||
|
default:
|
||||||
|
return LifespanData{}, fmt.Errorf("unknown lifespan type %q", attrs["type"])
|
||||||
|
}
|
||||||
|
|
||||||
|
valStr := attrs["val"]
|
||||||
|
totalStr := attrs["total"]
|
||||||
|
val, err := strconv.Atoi(valStr)
|
||||||
|
if err != nil {
|
||||||
|
return LifespanData{}, fmt.Errorf("invalid lifespan val %q", valStr)
|
||||||
|
}
|
||||||
|
total, err := strconv.Atoi(totalStr)
|
||||||
|
if err != nil {
|
||||||
|
return LifespanData{}, fmt.Errorf("invalid lifespan total %q", totalStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
return LifespanData{
|
||||||
|
Target: target,
|
||||||
|
Val: val,
|
||||||
|
Total: total,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/xml"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
_ "time/tzdata" // zone database for SetTime inside the distroless image
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
)
|
||||||
|
|
||||||
|
// readyCommand pairs one outbound ctl body with the command name registered
|
||||||
|
// in the correlator and surfaced in traces.
|
||||||
|
type readyCommand struct {
|
||||||
|
name string
|
||||||
|
out ctl.Outbound
|
||||||
|
}
|
||||||
|
|
||||||
|
// readySequence is the §8 fan-out: SetTime, the five Gets, then GetLifeSpan
|
||||||
|
// for SideBrush, Brush and DustCaseHeap. SetTime carries the current Unix
|
||||||
|
// second and the container-local UTC offset taken from now; Go division
|
||||||
|
// truncates toward zero so a negative offset keeps its sign on both parts.
|
||||||
|
func readySequence(now time.Time) []readyCommand {
|
||||||
|
_, offset := now.Zone()
|
||||||
|
hours := offset / 3600
|
||||||
|
minutes := (offset % 3600) / 60
|
||||||
|
|
||||||
|
seq := []readyCommand{{
|
||||||
|
name: "SetTime",
|
||||||
|
out: ctl.Outbound{
|
||||||
|
TD: "SetTime",
|
||||||
|
Inner: []byte(fmt.Sprintf(`<time t="%d" tz="%d" tzm="%d"/>`, now.Unix(), hours, minutes)),
|
||||||
|
},
|
||||||
|
}}
|
||||||
|
for _, td := range []string{"GetBatteryInfo", "GetCleanState", "GetChargeState", "GetCleanSpeed", "GetSched"} {
|
||||||
|
seq = append(seq, readyCommand{name: td, out: ctl.Outbound{TD: td}})
|
||||||
|
}
|
||||||
|
for _, kind := range []string{"SideBrush", "Brush", "DustCaseHeap"} {
|
||||||
|
seq = append(seq, readyCommand{
|
||||||
|
name: "GetLifeSpan",
|
||||||
|
out: ctl.Outbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: []ctl.Attr{{Name: "type", Value: kind}},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return seq
|
||||||
|
}
|
||||||
|
|
||||||
|
func cleanOutbound(cleanType, act, fan string) ctl.Outbound {
|
||||||
|
return ctl.Outbound{
|
||||||
|
TD: "Clean",
|
||||||
|
Inner: []byte(fmt.Sprintf(`<clean type="%s" speed="%s" act="%s"/>`,
|
||||||
|
escAttr(cleanType), escAttr(fan), escAttr(act))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func chargeOutbound(kind string) ctl.Outbound {
|
||||||
|
return ctl.Outbound{
|
||||||
|
TD: "Charge",
|
||||||
|
Inner: []byte(fmt.Sprintf(`<charge type="%s"/>`, escAttr(kind))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func playSoundOutbound() ctl.Outbound {
|
||||||
|
return ctl.Outbound{
|
||||||
|
TD: "PlaySound",
|
||||||
|
Attrs: []ctl.Attr{{Name: "sid", Value: "0"}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCleanSpeedOutbound(speed string) ctl.Outbound {
|
||||||
|
return ctl.Outbound{
|
||||||
|
TD: "SetCleanSpeed",
|
||||||
|
Attrs: []ctl.Attr{{Name: "speed", Value: speed}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// moveOutbound encodes Move without a ctl id; phase 05 originates it. The
|
||||||
|
// codec and its sid-only correlation are exercised by this phase's tests.
|
||||||
|
func moveOutbound(action string) ctl.Outbound {
|
||||||
|
return ctl.Outbound{
|
||||||
|
TD: "Move",
|
||||||
|
OmitCtlID: true,
|
||||||
|
Inner: []byte(fmt.Sprintf(`<move action="%s"/>`, escAttr(action))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func escAttr(v string) string {
|
||||||
|
var b strings.Builder
|
||||||
|
_ = xml.EscapeText(&b, []byte(v))
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
@@ -0,0 +1,965 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
testBotJID = "e20123456789@155.ecorobot.net/atom"
|
||||||
|
testSerial = "e20123456789"
|
||||||
|
testCtlJID = "n95bridge@ecouser.net/homeassistant"
|
||||||
|
)
|
||||||
|
|
||||||
|
// --- fakes and recorders ---
|
||||||
|
|
||||||
|
type fakeClock struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
now time.Time
|
||||||
|
timers []*fakeTimer
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeTimer struct {
|
||||||
|
at time.Time
|
||||||
|
ch chan time.Time
|
||||||
|
fired bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeClock() *fakeClock {
|
||||||
|
return &fakeClock{now: time.Unix(1790194386, 0)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeClock) Now() time.Time {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
return c.now
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeClock) After(d time.Duration) <-chan time.Time {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
t := &fakeTimer{at: c.now.Add(d), ch: make(chan time.Time, 1)}
|
||||||
|
c.timers = append(c.timers, t)
|
||||||
|
return t.ch
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeClock) advance(d time.Duration) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
c.now = c.now.Add(d)
|
||||||
|
for _, t := range c.timers {
|
||||||
|
if !t.fired && !t.at.After(c.now) {
|
||||||
|
t.fired = true
|
||||||
|
t.ch <- c.now
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type sendRec struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
xml []string
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *sendRec) send(b []byte) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.xml = append(s.xml, string(b))
|
||||||
|
return s.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *sendRec) all() []string {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
return append([]string(nil), s.xml...)
|
||||||
|
}
|
||||||
|
|
||||||
|
type pubRec struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
snaps []Snapshot
|
||||||
|
ch chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newPubRec() *pubRec { return &pubRec{ch: make(chan struct{}, 128)} }
|
||||||
|
|
||||||
|
func (r *pubRec) pub(_ context.Context, s Snapshot) {
|
||||||
|
r.mu.Lock()
|
||||||
|
r.snaps = append(r.snaps, s)
|
||||||
|
r.mu.Unlock()
|
||||||
|
r.ch <- struct{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *pubRec) last() Snapshot {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
if len(r.snaps) == 0 {
|
||||||
|
return Snapshot{}
|
||||||
|
}
|
||||||
|
return r.snaps[len(r.snaps)-1]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *pubRec) count() int {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return len(r.snaps)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *pubRec) wait(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
select {
|
||||||
|
case <-r.ch:
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for republish")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type traceRec struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
traces []ctl.Trace
|
||||||
|
ch chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTraceRec() *traceRec { return &traceRec{ch: make(chan struct{}, 128)} }
|
||||||
|
|
||||||
|
func (r *traceRec) sink(_ string, tr ctl.Trace) {
|
||||||
|
r.mu.Lock()
|
||||||
|
r.traces = append(r.traces, tr)
|
||||||
|
r.mu.Unlock()
|
||||||
|
r.ch <- struct{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *traceRec) wait(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
select {
|
||||||
|
case <-r.ch:
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for trace")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *traceRec) all() []ctl.Trace {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return append([]ctl.Trace(nil), r.traces...)
|
||||||
|
}
|
||||||
|
|
||||||
|
type diagRec struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
list []session.Diagnostic
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *diagRec) sink(d session.Diagnostic) {
|
||||||
|
r.mu.Lock()
|
||||||
|
r.list = append(r.list, d)
|
||||||
|
r.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *diagRec) all() []session.Diagnostic {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return append([]session.Diagnostic(nil), r.list...)
|
||||||
|
}
|
||||||
|
|
||||||
|
type rig struct {
|
||||||
|
actor *Actor
|
||||||
|
send *sendRec
|
||||||
|
pub *pubRec
|
||||||
|
trace *traceRec
|
||||||
|
diag *diagRec
|
||||||
|
clock *fakeClock
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRig(t *testing.T) *rig {
|
||||||
|
t.Helper()
|
||||||
|
send := &sendRec{}
|
||||||
|
pub := newPubRec()
|
||||||
|
tr := newTraceRec()
|
||||||
|
dg := &diagRec{}
|
||||||
|
clk := newFakeClock()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
t.Cleanup(cancel)
|
||||||
|
a := NewActor(ctx, testBotJID, testSerial, testCtlJID, pub.pub, tr.sink, dg.sink, clk)
|
||||||
|
a.SessionReady(session.ReadyEvent{
|
||||||
|
Generation: 1,
|
||||||
|
JID: testBotJID,
|
||||||
|
Serial: testSerial,
|
||||||
|
Send: send.send,
|
||||||
|
})
|
||||||
|
flushActor(t, a)
|
||||||
|
return &rig{actor: a, send: send, pub: pub, trace: tr, diag: dg, clock: clk}
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushActor blocks until every event queued so far has been processed.
|
||||||
|
func flushActor(t *testing.T, a *Actor) {
|
||||||
|
t.Helper()
|
||||||
|
ch := make(chan struct{})
|
||||||
|
select {
|
||||||
|
case a.mailbox <- evBarrier(ch):
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("actor mailbox blocked")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-ch:
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("actor did not process barrier")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *rig) stanza(t *testing.T, xml string) {
|
||||||
|
t.Helper()
|
||||||
|
r.actor.Stanza(session.StanzaEvent{
|
||||||
|
Generation: 1,
|
||||||
|
JID: testBotJID,
|
||||||
|
Serial: testSerial,
|
||||||
|
Stanza: []byte(xml),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- stanza builders ---
|
||||||
|
|
||||||
|
func iqResult(sid string) string {
|
||||||
|
return fmt.Sprintf(`<iq to="%s" type="result" id="%s"/>`, testCtlJID, sid)
|
||||||
|
}
|
||||||
|
|
||||||
|
func ctlResult(cid, ret, inner string) string {
|
||||||
|
return fmt.Sprintf(`<iq to="%s" type="set" id="9"><query xmlns="com:ctl"><ctl id="%s" ret="%s">%s</ctl></query></iq>`,
|
||||||
|
testCtlJID, cid, ret, inner)
|
||||||
|
}
|
||||||
|
|
||||||
|
func iqSet(frag string) string {
|
||||||
|
return `<iq to="` + testCtlJID + `" type="set" id="9"><query xmlns="com:ctl">` + frag + `</query></iq>`
|
||||||
|
}
|
||||||
|
|
||||||
|
func stanzaSID(t *testing.T, stanza string) string {
|
||||||
|
t.Helper()
|
||||||
|
i := strings.Index(stanza, `id="`)
|
||||||
|
if i < 0 {
|
||||||
|
t.Fatalf("no iq id in %s", stanza)
|
||||||
|
}
|
||||||
|
rest := stanza[i+4:]
|
||||||
|
j := strings.IndexByte(rest, '"')
|
||||||
|
return rest[:j]
|
||||||
|
}
|
||||||
|
|
||||||
|
func stanzaCID(t *testing.T, stanza string) string {
|
||||||
|
t.Helper()
|
||||||
|
in, err := ctl.Parse([]byte(stanza))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parse outbound: %v", err)
|
||||||
|
}
|
||||||
|
cid := in.Attrs["id"]
|
||||||
|
if cid == "" {
|
||||||
|
t.Fatalf("no ctl id in %s", stanza)
|
||||||
|
}
|
||||||
|
return cid
|
||||||
|
}
|
||||||
|
|
||||||
|
func applyPush(t *testing.T, snap *Snapshot, frag string) {
|
||||||
|
t.Helper()
|
||||||
|
in, err := ctl.Parse([]byte(iqSet(frag)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parse push: %v", err)
|
||||||
|
}
|
||||||
|
if err := Apply(snap, in, ""); err != nil {
|
||||||
|
t.Fatalf("Apply(%s): %v", frag, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- named phase 03 tests ---
|
||||||
|
|
||||||
|
func TestCtlSidVersusCid(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
r.actor.AnnounceOK(session.ReadyEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Send: r.send.send})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "start"}); err != nil {
|
||||||
|
t.Fatalf("Submit start: %v", err)
|
||||||
|
}
|
||||||
|
sends := r.send.all()
|
||||||
|
clean := sends[len(sends)-1]
|
||||||
|
sid, cid := stanzaSID(t, clean), stanzaCID(t, clean)
|
||||||
|
if sid == cid {
|
||||||
|
t.Fatalf("iq id %q must differ from ctl id", sid)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The stanza ack completes the ack phase only.
|
||||||
|
r.stanza(t, iqResult(sid))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
traces := r.trace.all()
|
||||||
|
if len(traces) != 1 || traces[0].Phase != "ack" || traces[0].SID != sid {
|
||||||
|
t.Fatalf("after ack, traces = %+v", traces)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got != nil {
|
||||||
|
t.Fatalf("ack must clear last_command_error, got %q", *got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A later iq-set whose ctl id matches completes the result phase once.
|
||||||
|
r.stanza(t, ctlResult(cid, "ok", ""))
|
||||||
|
r.stanza(t, ctlResult(cid, "ok", ""))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
traces = r.trace.all()
|
||||||
|
if len(traces) != 2 || traces[1].Phase != "result" || traces[1].Ret != "ok" {
|
||||||
|
t.Fatalf("after result, traces = %+v", traces)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetTime result without errno is success.
|
||||||
|
setTimeCID := stanzaCID(t, sends[0])
|
||||||
|
r.stanza(t, ctlResult(setTimeCID, "ok", ""))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got != nil {
|
||||||
|
t.Fatalf("ret ok without errno must clear last_command_error, got %q", *got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ret=fail sets last_command_error and leaves last_error null.
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "start"}); err != nil {
|
||||||
|
t.Fatalf("Submit start: %v", err)
|
||||||
|
}
|
||||||
|
sends = r.send.all()
|
||||||
|
failCID := stanzaCID(t, sends[len(sends)-1])
|
||||||
|
r.stanza(t, fmt.Sprintf(`<iq to="%s" type="set" id="9"><query xmlns="com:ctl"><ctl id="%s" ret="fail" errno="5"/></query></iq>`, testCtlJID, failCID))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
attrs := r.pub.last().Attributes
|
||||||
|
if got := attrs.LastCommandError; got == nil || *got != "fail:5" {
|
||||||
|
t.Fatalf("last_command_error = %v, want fail:5", got)
|
||||||
|
}
|
||||||
|
if attrs.LastError != nil {
|
||||||
|
t.Fatalf("last_error must stay null, got %q", *attrs.LastError)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadyFanOutOrder(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
if n := len(r.send.all()); n != 0 {
|
||||||
|
t.Fatalf("%d sends before AnnounceOK, want 0", n)
|
||||||
|
}
|
||||||
|
r.actor.AnnounceOK(session.ReadyEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Send: r.send.send})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
wantTD := []string{"SetTime", "GetBatteryInfo", "GetCleanState", "GetChargeState", "GetCleanSpeed", "GetSched", "GetLifeSpan", "GetLifeSpan", "GetLifeSpan"}
|
||||||
|
wantType := []string{"", "", "", "", "", "", "SideBrush", "Brush", "DustCaseHeap"}
|
||||||
|
sends := r.send.all()
|
||||||
|
if len(sends) != len(wantTD) {
|
||||||
|
t.Fatalf("%d fan-out stanzas, want %d", len(sends), len(wantTD))
|
||||||
|
}
|
||||||
|
seen := map[string]bool{}
|
||||||
|
for i, s := range sends {
|
||||||
|
in, err := ctl.Parse([]byte(s))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("stanza %d parse: %v", i, err)
|
||||||
|
}
|
||||||
|
if in.TD != wantTD[i] {
|
||||||
|
t.Errorf("stanza %d td = %q, want %q", i, in.TD, wantTD[i])
|
||||||
|
}
|
||||||
|
if wantType[i] != "" && in.Attrs["type"] != wantType[i] {
|
||||||
|
t.Errorf("stanza %d type = %q, want %q", i, in.Attrs["type"], wantType[i])
|
||||||
|
}
|
||||||
|
cid := in.Attrs["id"]
|
||||||
|
if cid == "" {
|
||||||
|
t.Errorf("stanza %d has no ctl id", i)
|
||||||
|
}
|
||||||
|
if seen[cid] {
|
||||||
|
t.Errorf("cid %q reused", cid)
|
||||||
|
}
|
||||||
|
seen[cid] = true
|
||||||
|
}
|
||||||
|
if !strings.Contains(sends[0], "<time t=") {
|
||||||
|
t.Errorf("SetTime missing time element: %s", sends[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStatePrecedenceTable(t *testing.T) {
|
||||||
|
e103 := "103"
|
||||||
|
rows := []struct {
|
||||||
|
name string
|
||||||
|
facts Facts
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"init", Facts{Fan: "standard"}, "idle"},
|
||||||
|
{"slot charging", Facts{Charge: "SlotCharging", Fan: "standard"}, "docked"},
|
||||||
|
{"going", Facts{Charge: "going", Fan: "standard"}, "returning"},
|
||||||
|
{"going beats clean auto", Facts{Charge: "going", CleanType: "auto", Fan: "standard"}, "returning"},
|
||||||
|
{"error latch", Facts{ErrorLatched: true, LastError: &e103, CleanType: "auto"}, "error"},
|
||||||
|
{"error beats docked stop", Facts{ErrorLatched: true, LastError: &e103, Charge: "SlotCharging", CleanType: "stop"}, "error"},
|
||||||
|
{"cleared error shows docked", Facts{Charge: "SlotCharging", CleanType: "stop"}, "docked"},
|
||||||
|
{"cleaning no charge", Facts{CleanType: "auto"}, "cleaning"},
|
||||||
|
{"idle charge with stop", Facts{Charge: "Idle", CleanType: "stop"}, "idle"},
|
||||||
|
{"idle charge with auto", Facts{Charge: "Idle", CleanType: "auto"}, "cleaning"},
|
||||||
|
{"going with stop", Facts{Charge: "going", CleanType: "stop"}, "returning"},
|
||||||
|
{"paused", Facts{Paused: true, CleanType: "auto"}, "paused"},
|
||||||
|
{"returning beats paused", Facts{Paused: true, Charge: "going"}, "returning"},
|
||||||
|
{"docked beats everything below", Facts{Charge: "SlotCharging", Paused: true, CleanType: "auto"}, "docked"},
|
||||||
|
}
|
||||||
|
for _, row := range rows {
|
||||||
|
if got := Derive(row.facts); got != row.want {
|
||||||
|
t.Errorf("%s: Derive = %q, want %q", row.name, got, row.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestErrno100ClearsLastErrorOnly(t *testing.T) {
|
||||||
|
snap := NewSnapshot()
|
||||||
|
applyPush(t, &snap, `<ctl td="ChargeState"><charge type="Idle"/></ctl>`)
|
||||||
|
applyPush(t, &snap, `<ctl td="CleanReport"><clean type="auto" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
if got := snap.State.State; got != "cleaning" {
|
||||||
|
t.Fatalf("state = %q, want cleaning", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
applyPush(t, &snap, `<ctl td="error" errno="103"/>`)
|
||||||
|
if got := snap.State.State; got != "error" {
|
||||||
|
t.Fatalf("state = %q, want error", got)
|
||||||
|
}
|
||||||
|
if le := snap.Attributes.LastError; le == nil || *le != "103" {
|
||||||
|
t.Fatalf("last_error = %v, want 103", le)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reports keep updating facts while the latch holds.
|
||||||
|
applyPush(t, &snap, `<ctl td="ChargeState"><charge type="SlotCharging"/></ctl>`)
|
||||||
|
applyPush(t, &snap, `<ctl td="CleanReport"><clean type="stop" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
if got := snap.Facts.CleanType; got != "stop" {
|
||||||
|
t.Fatalf("clean_type = %q, want stop", got)
|
||||||
|
}
|
||||||
|
if got := snap.State.State; got != "error" {
|
||||||
|
t.Fatalf("state = %q, want error while latched", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
applyPush(t, &snap, `<ctl td="error" errno="100"/>`)
|
||||||
|
if snap.Attributes.LastError != nil {
|
||||||
|
t.Fatalf("last_error = %v, want null after errno 100", *snap.Attributes.LastError)
|
||||||
|
}
|
||||||
|
if got := snap.State.State; got != "docked" {
|
||||||
|
t.Fatalf("state = %q, want docked from SlotCharging + stop", got)
|
||||||
|
}
|
||||||
|
if got := snap.Facts.Charge; got != "SlotCharging" {
|
||||||
|
t.Fatalf("charge = %q, errno 100 must not clear facts", got)
|
||||||
|
}
|
||||||
|
if got := snap.Facts.Fan; got != "standard" {
|
||||||
|
t.Fatalf("fan = %q, errno 100 must not clear fan", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdleReDerivesWithoutStartingClean(t *testing.T) {
|
||||||
|
snap := NewSnapshot()
|
||||||
|
applyPush(t, &snap, `<ctl td="ChargeState"><charge type="SlotCharging"/></ctl>`)
|
||||||
|
applyPush(t, &snap, `<ctl td="CleanReport"><clean type="stop" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
applyPush(t, &snap, `<ctl td="ChargeState"><charge type="Idle"/></ctl>`)
|
||||||
|
if got := snap.State.State; got != "idle" {
|
||||||
|
t.Fatalf("Idle + stop: state = %q, want idle", got)
|
||||||
|
}
|
||||||
|
if got := snap.Facts.CleanType; got != "stop" {
|
||||||
|
t.Fatalf("Idle must not write clean_type, got %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
snap2 := NewSnapshot()
|
||||||
|
applyPush(t, &snap2, `<ctl td="ChargeState"><charge type="SlotCharging"/></ctl>`)
|
||||||
|
applyPush(t, &snap2, `<ctl td="CleanReport"><clean type="auto" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
applyPush(t, &snap2, `<ctl td="ChargeState"><charge type="Idle"/></ctl>`)
|
||||||
|
if got := snap2.State.State; got != "cleaning" {
|
||||||
|
t.Fatalf("Idle + remembered auto: state = %q, want cleaning", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A direct td="Idle" push is the same report: it rewrites charge_state
|
||||||
|
// and re-derives without writing a clean type.
|
||||||
|
snap3 := NewSnapshot()
|
||||||
|
applyPush(t, &snap3, `<ctl td="ChargeState"><charge type="SlotCharging"/></ctl>`)
|
||||||
|
applyPush(t, &snap3, `<ctl td="CleanReport"><clean type="stop" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
applyPush(t, &snap3, `<ctl td="Idle"/>`)
|
||||||
|
if got := snap3.State.State; got != "idle" {
|
||||||
|
t.Fatalf("td=Idle + stop: state = %q, want idle", got)
|
||||||
|
}
|
||||||
|
if got := snap3.Facts.CleanType; got != "stop" {
|
||||||
|
t.Fatalf("td=Idle must not write clean_type, got %q", got)
|
||||||
|
}
|
||||||
|
if got := snap3.Facts.Charge; got != "Idle" {
|
||||||
|
t.Fatalf("td=Idle must set charge_state, got %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
snap4 := NewSnapshot()
|
||||||
|
applyPush(t, &snap4, `<ctl td="ChargeState"><charge type="SlotCharging"/></ctl>`)
|
||||||
|
applyPush(t, &snap4, `<ctl td="CleanReport"><clean type="auto" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
applyPush(t, &snap4, `<ctl td="Idle"/>`)
|
||||||
|
if got := snap4.State.State; got != "cleaning" {
|
||||||
|
t.Fatalf("td=Idle + remembered auto: state = %q, want cleaning", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStopDockedOnlyWhileSlotCharging(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
charge string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"SlotCharging", "docked"},
|
||||||
|
{"Idle", "idle"},
|
||||||
|
{"going", "returning"},
|
||||||
|
} {
|
||||||
|
snap := NewSnapshot()
|
||||||
|
applyPush(t, &snap, `<ctl td="ChargeState"><charge type="`+tc.charge+`"/></ctl>`)
|
||||||
|
applyPush(t, &snap, `<ctl td="CleanReport"><clean type="stop" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
if got := snap.State.State; got != tc.want {
|
||||||
|
t.Errorf("stop + %s: state = %q, want %q", tc.charge, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSetCleanSpeedStoresRequested(t *testing.T) {
|
||||||
|
snap := NewSnapshot()
|
||||||
|
|
||||||
|
// A SetCleanSpeed result ret="ok" errno="" has no speed echo; the
|
||||||
|
// requested fan is stored anyway.
|
||||||
|
in := ctl.Inbound{Kind: ctl.KindResult, TD: "SetCleanSpeed", Ret: "ok", Errno: strPtr(""), Attrs: map[string]string{}}
|
||||||
|
if err := Apply(&snap, in, "strong"); err != nil {
|
||||||
|
t.Fatalf("Apply SetCleanSpeed result: %v", err)
|
||||||
|
}
|
||||||
|
if snap.Facts.Fan != "strong" || snap.State.FanSpeed != "strong" {
|
||||||
|
t.Fatalf("fan = %q/%q, want strong", snap.Facts.Fan, snap.State.FanSpeed)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A later CleanReport at standard replaces it.
|
||||||
|
applyPush(t, &snap, `<ctl td="CleanReport"><clean type="auto" speed="standard" st=" " rsn=" "/></ctl>`)
|
||||||
|
if snap.Facts.Fan != "standard" {
|
||||||
|
t.Fatalf("fan = %q, want standard from CleanReport", snap.Facts.Fan)
|
||||||
|
}
|
||||||
|
|
||||||
|
// An unknown speed is rejected and leaves the fan untouched.
|
||||||
|
bad := ctl.Inbound{Kind: ctl.KindPush, TD: "CleanReport", CleanAttrs: map[string]string{"type": "auto", "speed": "ludicrous"}}
|
||||||
|
if err := Apply(&snap, bad, ""); err == nil {
|
||||||
|
t.Fatal("Apply with unknown speed returned nil error")
|
||||||
|
}
|
||||||
|
if snap.Facts.Fan != "standard" {
|
||||||
|
t.Fatalf("fan = %q after invalid speed, want standard", snap.Facts.Fan)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBareBatteryNotAckedAndCidCompletesOnce(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
r.stanza(t, `<iq to="`+testCtlJID+`" type="set" id="46"><query xmlns="com:ctl"><battery power="076"/></query></iq>`)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
if n := len(r.send.all()); n != 0 {
|
||||||
|
t.Fatalf("bare battery produced %d outbound stanzas, want 0", n)
|
||||||
|
}
|
||||||
|
if b := r.pub.last().Attributes.BatteryLevel; b == nil || *b != 76 {
|
||||||
|
t.Fatalf("battery_level = %v, want 76", b)
|
||||||
|
}
|
||||||
|
|
||||||
|
c := ctl.NewCorrelator(nil)
|
||||||
|
c.Register("1", "00000042", "GetCleanState", true)
|
||||||
|
c.Register("2", "00000042", "GetCleanState", true)
|
||||||
|
if _, ok := c.CompleteCID("00000042"); !ok {
|
||||||
|
t.Fatal("first cid completion failed")
|
||||||
|
}
|
||||||
|
if _, ok := c.CompleteCID("00000042"); ok {
|
||||||
|
t.Fatal("cid completed twice")
|
||||||
|
}
|
||||||
|
c.Register("3", "00000042", "GetCleanState", true)
|
||||||
|
if _, ok := c.CompleteCID("00000042"); !ok {
|
||||||
|
t.Fatal("cid was not reusable after completion")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoAckForUnsolicitedIQSet(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
r.actor.AnnounceOK(session.ReadyEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Send: r.send.send})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
r.stanza(t, iqSet(`<ctl td="CleanReport"><clean type="auto" speed="standard" st=" " rsn=" "/></ctl>`))
|
||||||
|
r.stanza(t, ctlResult(stanzaCID(t, r.send.all()[0]), "ok", ""))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
for _, s := range r.send.all() {
|
||||||
|
if strings.Contains(s, `type="result"`) {
|
||||||
|
t.Fatalf("actor sent an iq result for an unsolicited iq-set: %s", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOutstandingCidFailedOnReplace(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "start"}); err != nil {
|
||||||
|
t.Fatalf("Submit: %v", err)
|
||||||
|
}
|
||||||
|
sends := r.send.all()
|
||||||
|
cid := stanzaCID(t, sends[len(sends)-1])
|
||||||
|
|
||||||
|
r.actor.SessionDown(session.DownEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Reason: session.ReasonReplaced})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
traces := r.trace.all()
|
||||||
|
if len(traces) != 1 || traces[0].Ret != "connection-lost" || traces[0].CID != cid {
|
||||||
|
t.Fatalf("traces after SessionDown = %+v", traces)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "connection-lost" {
|
||||||
|
t.Fatalf("last_command_error = %v, want connection-lost", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A new generation does not resurrect the old cid: a late result for it
|
||||||
|
// is ignored.
|
||||||
|
r.actor.SessionReady(session.ReadyEvent{Generation: 2, JID: testBotJID, Serial: testSerial, Send: r.send.send})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
r.stanza(t, ctlResult(cid, "ok", ""))
|
||||||
|
r.actor.Stanza(session.StanzaEvent{Generation: 2, JID: testBotJID, Serial: testSerial, Stanza: []byte(ctlResult(cid, "ok", ""))})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
if n := len(r.trace.all()); n != 1 {
|
||||||
|
t.Fatalf("late result produced %d traces total, want 1", n)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "connection-lost" {
|
||||||
|
t.Fatalf("last_command_error = %v, want connection-lost still", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCommandTimeoutExpiresCorrelation(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "start"}); err != nil {
|
||||||
|
t.Fatalf("Submit: %v", err)
|
||||||
|
}
|
||||||
|
r.clock.advance(ctl.CommandTimeout + time.Second)
|
||||||
|
r.trace.wait(t)
|
||||||
|
|
||||||
|
tr := r.trace.all()[0]
|
||||||
|
if tr.Ret != "timeout" || tr.Phase != "result" {
|
||||||
|
t.Fatalf("expire trace = %+v, want timeout/result", tr)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "timeout" {
|
||||||
|
t.Fatalf("last_command_error = %v, want timeout", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A command registered 29s after actor start must expire at its own 30-second
|
||||||
|
// deadline, not at the next coarse timer tick.
|
||||||
|
func TestCorrelationExpiresAtOwnDeadline(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
r.clock.advance(29 * time.Second)
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "start"}); err != nil {
|
||||||
|
t.Fatalf("Submit: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 29s after registration: not yet expired.
|
||||||
|
r.clock.advance(29 * time.Second)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
if traces := r.trace.all(); len(traces) != 0 {
|
||||||
|
t.Fatalf("command expired early: %+v", traces)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Registration + 30s: timeout fires on the command's own deadline.
|
||||||
|
r.clock.advance(time.Second)
|
||||||
|
r.trace.wait(t)
|
||||||
|
tr := r.trace.all()[0]
|
||||||
|
if tr.Ret != "timeout" || tr.Phase != "result" {
|
||||||
|
t.Fatalf("expire trace = %+v, want timeout/result", tr)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "timeout" {
|
||||||
|
t.Fatalf("last_command_error = %v, want timeout", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A SetCleanSpeed result ret="ok" with no speed echo stores the requested fan,
|
||||||
|
// and the cid entry is consumed. This drives the actor path end to end; the
|
||||||
|
// Apply-level case lives in TestSetCleanSpeedStoresRequested.
|
||||||
|
func TestSetCleanSpeedResultStoresFan(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "set_fan_speed", Args: map[string]string{"speed": "strong"}}); err != nil {
|
||||||
|
t.Fatalf("Submit: %v", err)
|
||||||
|
}
|
||||||
|
sends := r.send.all()
|
||||||
|
cid := stanzaCID(t, sends[len(sends)-1])
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
if got := r.actor.requestedFan[cid]; got != "strong" {
|
||||||
|
t.Fatalf("requestedFan[%s] = %q, want strong", cid, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.stanza(t, fmt.Sprintf(`<iq to="%s" type="set" id="9"><query xmlns="com:ctl"><ctl id="%s" ret="ok" errno=""/></query></iq>`, testCtlJID, cid))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
if got := r.pub.last().Facts.Fan; got != "strong" {
|
||||||
|
t.Fatalf("fan = %q, want strong", got)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().State.FanSpeed; got != "strong" {
|
||||||
|
t.Fatalf("fan_speed = %q, want strong", got)
|
||||||
|
}
|
||||||
|
if _, ok := r.actor.requestedFan[cid]; ok {
|
||||||
|
t.Fatal("requestedFan not cleared after result")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A SetCleanSpeed result that never arrives must not pin the requested fan:
|
||||||
|
// the timeout clears the cid entry and leaves the fan unchanged.
|
||||||
|
func TestSetCleanSpeedExpireClearsRequestedFan(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "set_fan_speed", Args: map[string]string{"speed": "strong"}}); err != nil {
|
||||||
|
t.Fatalf("Submit: %v", err)
|
||||||
|
}
|
||||||
|
sends := r.send.all()
|
||||||
|
cid := stanzaCID(t, sends[len(sends)-1])
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
if got := r.actor.requestedFan[cid]; got != "strong" {
|
||||||
|
t.Fatalf("requestedFan[%s] = %q, want strong", cid, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.clock.advance(ctl.CommandTimeout + time.Second)
|
||||||
|
r.trace.wait(t)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
if _, ok := r.actor.requestedFan[cid]; ok {
|
||||||
|
t.Fatal("requestedFan retained after timeout")
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "timeout" {
|
||||||
|
t.Fatalf("last_command_error = %v, want timeout", got)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Facts.Fan; got != "standard" {
|
||||||
|
t.Fatalf("fan = %q, want unchanged standard", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSubmitOfflineAndRejected(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
// Drop the generation so the actor is offline.
|
||||||
|
r.actor.SessionDown(session.DownEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Reason: session.ReasonTCPClose})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "start"}); err == nil || err.Error() != "offline" {
|
||||||
|
t.Fatalf("offline Submit err = %v, want offline", err)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "offline" {
|
||||||
|
t.Fatalf("last_command_error = %v, want offline", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unknown names and bad fan speeds are rejected without a write.
|
||||||
|
r.actor.SessionReady(session.ReadyEvent{Generation: 2, JID: testBotJID, Serial: testSerial, Send: r.send.send})
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
before := len(r.send.all())
|
||||||
|
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "dance"}); err == nil || err.Error() != "rejected:dance" {
|
||||||
|
t.Fatalf("unknown Submit err = %v, want rejected:dance", err)
|
||||||
|
}
|
||||||
|
if err := r.actor.Submit(context.Background(), Command{Name: "set_fan_speed", Args: map[string]string{"speed": "turbo"}}); err == nil || err.Error() != "rejected:set_fan_speed" {
|
||||||
|
t.Fatalf("bad speed Submit err = %v, want rejected:set_fan_speed", err)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "rejected:set_fan_speed" {
|
||||||
|
t.Fatalf("last_command_error = %v, want rejected:set_fan_speed", got)
|
||||||
|
}
|
||||||
|
if n := len(r.send.all()); n != before {
|
||||||
|
t.Fatalf("rejected commands wrote %d stanzas", n-before)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStaleSendFailsCommandOnce(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
r.send.err = session.ErrStale
|
||||||
|
|
||||||
|
err := r.actor.Submit(context.Background(), Command{Name: "start"})
|
||||||
|
if err == nil || err.Error() != "connection-lost" {
|
||||||
|
t.Fatalf("stale Submit err = %v, want connection-lost", err)
|
||||||
|
}
|
||||||
|
// One write attempt, no retry on the stale generation.
|
||||||
|
if n := len(r.send.all()); n != 1 {
|
||||||
|
t.Fatalf("%d write attempts, want 1", n)
|
||||||
|
}
|
||||||
|
traces := r.trace.all()
|
||||||
|
if len(traces) != 1 || traces[0].Ret != "connection-lost" {
|
||||||
|
t.Fatalf("traces = %+v, want one connection-lost", traces)
|
||||||
|
}
|
||||||
|
if got := r.pub.last().Attributes.LastCommandError; got == nil || *got != "connection-lost" {
|
||||||
|
t.Fatalf("last_command_error = %v, want connection-lost", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRepublishFullObjects(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
r.stanza(t, iqSet(`<ctl td="BatteryInfo"><battery power="076"/></ctl>`))
|
||||||
|
r.pub.wait(t)
|
||||||
|
|
||||||
|
snap := r.pub.last()
|
||||||
|
stateJSON, err := json.Marshal(snap.State)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal state: %v", err)
|
||||||
|
}
|
||||||
|
var state map[string]any
|
||||||
|
if err := json.Unmarshal(stateJSON, &state); err != nil {
|
||||||
|
t.Fatalf("unmarshal state: %v", err)
|
||||||
|
}
|
||||||
|
if state["state"] != "idle" || state["fan_speed"] != "standard" {
|
||||||
|
t.Fatalf("state doc = %s", stateJSON)
|
||||||
|
}
|
||||||
|
if len(state) != 2 {
|
||||||
|
t.Fatalf("state doc has %d keys, want state and fan_speed: %s", len(state), stateJSON)
|
||||||
|
}
|
||||||
|
|
||||||
|
attrJSON, err := json.Marshal(snap.Attributes)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal attributes: %v", err)
|
||||||
|
}
|
||||||
|
var attrs map[string]any
|
||||||
|
if err := json.Unmarshal(attrJSON, &attrs); err != nil {
|
||||||
|
t.Fatalf("unmarshal attributes: %v", err)
|
||||||
|
}
|
||||||
|
for _, key := range []string{
|
||||||
|
"battery_level", "side_brush", "main_brush", "filter",
|
||||||
|
"lifespan_total", "clean_type", "charge_state",
|
||||||
|
"last_error", "last_command_error", "schedules",
|
||||||
|
} {
|
||||||
|
if _, ok := attrs[key]; !ok {
|
||||||
|
t.Errorf("attributes missing %q: %s", key, attrJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if attrs["battery_level"] != float64(76) {
|
||||||
|
t.Errorf("battery_level = %v, want 76", attrs["battery_level"])
|
||||||
|
}
|
||||||
|
if !strings.Contains(string(attrJSON), `"schedules":[]`) {
|
||||||
|
t.Errorf("schedules must marshal as []: %s", attrJSON)
|
||||||
|
}
|
||||||
|
lt, ok := attrs["lifespan_total"].(map[string]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("lifespan_total = %v, want object", attrs["lifespan_total"])
|
||||||
|
}
|
||||||
|
for _, key := range []string{"side_brush", "main_brush", "filter"} {
|
||||||
|
if v, ok := lt[key]; !ok || v != nil {
|
||||||
|
t.Errorf("lifespan_total.%s = %v, want null", key, v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// requireDiag finds the diagnostic emitted for an exact inbound stanza whose
|
||||||
|
// reason contains want.
|
||||||
|
func requireDiag(t *testing.T, dg *diagRec, xml, want string) session.Diagnostic {
|
||||||
|
t.Helper()
|
||||||
|
for _, d := range dg.all() {
|
||||||
|
if string(d.XML) == xml && strings.Contains(d.Reason, want) {
|
||||||
|
if d.Kind != "unparsed" || d.Direction != session.DirectionIn {
|
||||||
|
t.Fatalf("diagnostic = %+v, want kind unparsed direction in", d)
|
||||||
|
}
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.Fatalf("no diagnostic for %q with reason containing %q in %+v", xml, want, dg.all())
|
||||||
|
return session.Diagnostic{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnparsedDiagnostics(t *testing.T) {
|
||||||
|
r := newRig(t)
|
||||||
|
|
||||||
|
// Unknown ctl td: diagnosed as unparsed, no mutation, no republish.
|
||||||
|
mystery := iqSet(`<ctl td="Mystery"><zap/></ctl>`)
|
||||||
|
r.stanza(t, mystery)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
requireDiag(t, r.diag, mystery, "Mystery")
|
||||||
|
if n := r.pub.count(); n != 0 {
|
||||||
|
t.Fatalf("unknown ctl produced %d republishes, want 0", n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Malformed XML.
|
||||||
|
malformed := `<iq type="set"><query><ctl td="x"`
|
||||||
|
r.stanza(t, malformed)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
requireDiag(t, r.diag, malformed, "malformed")
|
||||||
|
|
||||||
|
// Well-formed stanza that is not a ctl shape.
|
||||||
|
unknownShape := `<iq type="get" id="5"><query/></iq>`
|
||||||
|
r.stanza(t, unknownShape)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
requireDiag(t, r.diag, unknownShape, "unparsed")
|
||||||
|
|
||||||
|
// Out-of-range battery: diagnosed, attribute untouched, no republish.
|
||||||
|
badBattery := iqSet(`<ctl td="BatteryInfo"><battery power="999"/></ctl>`)
|
||||||
|
pubs := r.pub.count()
|
||||||
|
r.stanza(t, badBattery)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
requireDiag(t, r.diag, badBattery, "999")
|
||||||
|
if n := r.pub.count(); n != pubs {
|
||||||
|
t.Fatalf("invalid battery produced a republish")
|
||||||
|
}
|
||||||
|
if b := r.pub.last().Attributes.BatteryLevel; b != nil {
|
||||||
|
t.Fatalf("battery_level = %v, want null after invalid report", *b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Nonblank CleanReport st/rsn: diagnosed with key/value, while the
|
||||||
|
// canonical clean type and fan still apply and republish.
|
||||||
|
report := iqSet(`<ctl td="CleanReport"><clean type="auto" speed="strong" st="x" rsn="wheels"/></ctl>`)
|
||||||
|
pubs = r.pub.count()
|
||||||
|
r.stanza(t, report)
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
d := requireDiag(t, r.diag, report, "st")
|
||||||
|
if !strings.Contains(d.Reason, "rsn") || !strings.Contains(d.Reason, "wheels") {
|
||||||
|
t.Fatalf("diagnostic reason %q missing rsn detail", d.Reason)
|
||||||
|
}
|
||||||
|
if n := r.pub.count(); n <= pubs {
|
||||||
|
t.Fatal("valid canonical fields were not republished")
|
||||||
|
}
|
||||||
|
snap := r.pub.last()
|
||||||
|
if got := snap.Facts.CleanType; got != "auto" {
|
||||||
|
t.Fatalf("clean_type = %q, want auto", got)
|
||||||
|
}
|
||||||
|
if got := snap.Facts.Fan; got != "strong" {
|
||||||
|
t.Fatalf("fan = %q, want strong", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The session is untouched throughout: a valid push still applies.
|
||||||
|
r.stanza(t, iqSet(`<ctl td="BatteryInfo"><battery power="050"/></ctl>`))
|
||||||
|
flushActor(t, r.actor)
|
||||||
|
if b := r.pub.last().Attributes.BatteryLevel; b == nil || *b != 50 {
|
||||||
|
t.Fatalf("battery_level = %v after diagnostics, want 50", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNonblankGetCleanStateFieldsDiagnose(t *testing.T) {
|
||||||
|
snap := NewSnapshot()
|
||||||
|
in := ctl.Inbound{
|
||||||
|
Kind: ctl.KindResult,
|
||||||
|
TD: "GetCleanState",
|
||||||
|
CleanAttrs: map[string]string{
|
||||||
|
"type": "stop", "speed": "standard", "st": "h", "t": "123", "a": " ",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
err := Apply(&snap, in, "")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("Apply with nonblank t returned nil error")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "t=") || !strings.Contains(err.Error(), "123") {
|
||||||
|
t.Fatalf("Apply error = %v, want t= key/value", err)
|
||||||
|
}
|
||||||
|
if strings.Contains(err.Error(), "a=") {
|
||||||
|
t.Fatalf("whitespace a must be ignored, got %v", err)
|
||||||
|
}
|
||||||
|
if snap.Facts.CleanType != "stop" {
|
||||||
|
t.Fatalf("clean_type = %q, canonical update must still apply", snap.Facts.CleanType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFleetCreatesAndForwards(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
pub := newPubRec()
|
||||||
|
send := &sendRec{}
|
||||||
|
f := NewFleet(ctx, testCtlJID, pub.pub, nil, nil, nil)
|
||||||
|
|
||||||
|
f.SessionReady(session.ReadyEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Send: send.send})
|
||||||
|
a, ok := f.Actor(testBotJID)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("Fleet did not create an actor on SessionReady")
|
||||||
|
}
|
||||||
|
if _, ok := f.Actor("other@155.ecorobot.net/atom"); ok {
|
||||||
|
t.Fatal("Fleet returned an actor for an unknown JID")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Events for the JID reach the actor; events for other JIDs are dropped.
|
||||||
|
f.Stanza(session.StanzaEvent{Generation: 1, JID: testBotJID, Serial: testSerial, Stanza: []byte(iqSet(`<ctl td="BatteryInfo"><battery power="060"/></ctl>`))})
|
||||||
|
flushActor(t, a)
|
||||||
|
if b := pub.last().Attributes.BatteryLevel; b == nil || *b != 60 {
|
||||||
|
t.Fatalf("battery_level = %v, want 60", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func strPtr(v string) *string { return &v }
|
||||||
@@ -0,0 +1,73 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/xml"
|
||||||
|
"errors"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type schedEntry struct {
|
||||||
|
Name string `xml:"n,attr"`
|
||||||
|
On string `xml:"o,attr"`
|
||||||
|
Time string `xml:"t,attr"`
|
||||||
|
Repeat string `xml:"r,attr"`
|
||||||
|
Flag string `xml:"f,attr"`
|
||||||
|
Action struct {
|
||||||
|
TD string `xml:"td,attr"`
|
||||||
|
Type string `xml:"type,attr"`
|
||||||
|
} `xml:"ctl"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseSchedulesImpl implements ScheduleParser.
|
||||||
|
func ScheduleParserHook(inner []byte) ([]Schedule, error) {
|
||||||
|
if len(bytes.TrimSpace(inner)) == 0 {
|
||||||
|
return []Schedule{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Because `inner` contains a sequence of `<s>...</s>` elements,
|
||||||
|
// we wrap it in a dummy root to use standard xml.Unmarshal.
|
||||||
|
wrapped := []byte("<root>")
|
||||||
|
wrapped = append(wrapped, inner...)
|
||||||
|
wrapped = append(wrapped, []byte("</root>")...)
|
||||||
|
|
||||||
|
var root struct {
|
||||||
|
Entries []schedEntry `xml:"s"`
|
||||||
|
}
|
||||||
|
if err := xml.Unmarshal(wrapped, &root); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var schedules []Schedule
|
||||||
|
var errs []error
|
||||||
|
for _, entry := range root.Entries {
|
||||||
|
if entry.Action.Type != "auto" {
|
||||||
|
errs = append(errs, errors.New("unsupported schedule type: "+entry.Action.Type))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
s := Schedule{
|
||||||
|
Name: entry.Name,
|
||||||
|
On: entry.On == "1",
|
||||||
|
Time: entry.Time,
|
||||||
|
Repeat: entry.Repeat,
|
||||||
|
Flag: entry.Flag,
|
||||||
|
Action: ScheduleAction{
|
||||||
|
TD: entry.Action.TD,
|
||||||
|
Type: entry.Action.Type,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
schedules = append(schedules, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
var err error
|
||||||
|
if len(errs) > 0 {
|
||||||
|
err = errors.Join(errs...)
|
||||||
|
}
|
||||||
|
return schedules, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// RepeatIndex returns the Sunday-first index for the given weekday.
|
||||||
|
func RepeatIndex(d time.Weekday) int {
|
||||||
|
return int(d)
|
||||||
|
}
|
||||||
@@ -0,0 +1,166 @@
|
|||||||
|
package robot_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
|
||||||
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
robot.ParseSchedules = robot.ScheduleParserHook
|
||||||
|
robot.ParseLifespan = robot.LifespanParserHook
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSchedulesSundayFirstAndReplace(t *testing.T) {
|
||||||
|
if got := robot.RepeatIndex(time.Sunday); got != 0 {
|
||||||
|
t.Errorf("Sunday = %d, want 0", got)
|
||||||
|
}
|
||||||
|
if got := robot.RepeatIndex(time.Wednesday); got != 3 {
|
||||||
|
t.Errorf("Wednesday = %d, want 3", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
|
||||||
|
in1 := ctl.Inbound{
|
||||||
|
TD: "Sched2",
|
||||||
|
Inner: []byte(`
|
||||||
|
<s n="1" o="1" t="10:00" r="0001000" f="p"><ctl td="Clean" type="auto"/></s>
|
||||||
|
`),
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in1, "")
|
||||||
|
if len(snap.Attributes.Schedules) != 1 {
|
||||||
|
t.Fatalf("len = %d, want 1", len(snap.Attributes.Schedules))
|
||||||
|
}
|
||||||
|
|
||||||
|
in2 := ctl.Inbound{
|
||||||
|
TD: "Sched2",
|
||||||
|
Inner: []byte(`
|
||||||
|
<s n="2" o="0" t="11:00" r="1111111" f="p"><ctl td="Clean" type="auto"/></s>
|
||||||
|
<s n="3" o="1" t="12:00" r="0000000" f="p"><ctl td="Clean" type="auto"/></s>
|
||||||
|
`),
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in2, "")
|
||||||
|
if len(snap.Attributes.Schedules) != 2 {
|
||||||
|
t.Fatalf("len = %d, want 2", len(snap.Attributes.Schedules))
|
||||||
|
}
|
||||||
|
if snap.Attributes.Schedules[0].Name != "2" || snap.Attributes.Schedules[1].Name != "3" {
|
||||||
|
t.Errorf("schedules not replaced entirely")
|
||||||
|
}
|
||||||
|
|
||||||
|
in3 := ctl.Inbound{
|
||||||
|
TD: "GetSched",
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in3, "")
|
||||||
|
if len(snap.Attributes.Schedules) != 0 {
|
||||||
|
t.Fatalf("len = %d, want 0", len(snap.Attributes.Schedules))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEmptySchedShapes(t *testing.T) {
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
|
||||||
|
in1 := ctl.Inbound{
|
||||||
|
TD: "Sched2",
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in1, "")
|
||||||
|
if snap.Attributes.Schedules == nil || len(snap.Attributes.Schedules) != 0 {
|
||||||
|
t.Errorf("Sched2 without inner should be []")
|
||||||
|
}
|
||||||
|
|
||||||
|
in2 := ctl.Inbound{
|
||||||
|
TD: "GetSched",
|
||||||
|
Inner: []byte{},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in2, "")
|
||||||
|
if snap.Attributes.Schedules == nil || len(snap.Attributes.Schedules) != 0 {
|
||||||
|
t.Errorf("GetSched without inner should be []")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestScheduleWhitespaceIgnored(t *testing.T) {
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
in := ctl.Inbound{
|
||||||
|
TD: "Sched2",
|
||||||
|
Inner: []byte(`<s n="1" o="1" t="10:00" r="0001000" f="p"> <ctl td="clean" type="auto"/> </s>`),
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in, "")
|
||||||
|
if len(snap.Attributes.Schedules) != 1 {
|
||||||
|
t.Fatalf("len = %d, want 1", len(snap.Attributes.Schedules))
|
||||||
|
}
|
||||||
|
s := snap.Attributes.Schedules[0]
|
||||||
|
if s.Flag != "p" || s.Action.Type != "auto" {
|
||||||
|
t.Errorf("bad parse: %+v", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestScheduleRejectsNonAuto(t *testing.T) {
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
in := ctl.Inbound{
|
||||||
|
TD: "Sched2",
|
||||||
|
Inner: []byte(`
|
||||||
|
<s n="1" o="1" t="10:00" r="0001000" f="p"><ctl td="Clean" type="border"/></s>
|
||||||
|
<s n="2" o="1" t="11:00" r="0001000" f="p"><ctl td="Clean" type="auto"/></s>
|
||||||
|
`),
|
||||||
|
}
|
||||||
|
err := robot.Apply(&snap, in, "")
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("expected error for non-auto schedule")
|
||||||
|
}
|
||||||
|
if len(snap.Attributes.Schedules) != 1 {
|
||||||
|
t.Fatalf("len = %d, want 1", len(snap.Attributes.Schedules))
|
||||||
|
}
|
||||||
|
if snap.Attributes.Schedules[0].Name != "2" {
|
||||||
|
t.Errorf("wrong schedule kept: %s", snap.Attributes.Schedules[0].Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLifespanUnknownUnit(t *testing.T) {
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
|
||||||
|
in1 := ctl.Inbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: map[string]string{"type": "SideBrush", "val": "068", "total": "365"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in1, "")
|
||||||
|
|
||||||
|
in2 := ctl.Inbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: map[string]string{"type": "Brush", "val": "90", "total": "365"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in2, "")
|
||||||
|
|
||||||
|
in3 := ctl.Inbound{
|
||||||
|
TD: "GetLifeSpan",
|
||||||
|
Attrs: map[string]string{"type": "DustCaseHeap", "val": "70", "total": "365"},
|
||||||
|
}
|
||||||
|
robot.Apply(&snap, in3, "")
|
||||||
|
|
||||||
|
a := snap.Attributes
|
||||||
|
if a.SideBrush == nil || *a.SideBrush != 68 {
|
||||||
|
t.Errorf("SideBrush = %v", a.SideBrush)
|
||||||
|
}
|
||||||
|
if a.LifespanTotal.SideBrush == nil || *a.LifespanTotal.SideBrush != 365 {
|
||||||
|
t.Errorf("SideBrush total = %v", a.LifespanTotal.SideBrush)
|
||||||
|
}
|
||||||
|
if a.MainBrush == nil || *a.MainBrush != 90 {
|
||||||
|
t.Errorf("MainBrush = %v", a.MainBrush)
|
||||||
|
}
|
||||||
|
if a.Filter == nil || *a.Filter != 70 {
|
||||||
|
t.Errorf("Filter = %v", a.Filter)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestScheduleCommandResultsAccepted(t *testing.T) {
|
||||||
|
snap := robot.NewSnapshot()
|
||||||
|
for _, td := range []string{"AddSched", "ModSched", "DelSched"} {
|
||||||
|
in := ctl.Inbound{
|
||||||
|
TD: td,
|
||||||
|
Ret: "ok",
|
||||||
|
}
|
||||||
|
if err := robot.Apply(&snap, in, ""); err != nil {
|
||||||
|
t.Errorf("Apply failed for %s result: %v", td, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,327 @@
|
|||||||
|
// Package robot owns the per-robot actor, the READY ctl fan-out, and the
|
||||||
|
// retained state and attribute documents from N95-FULL-SPECIFICATION.md §10.
|
||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/ctl"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Facts are the retained inputs the Derive precedence table reads.
|
||||||
|
type Facts struct {
|
||||||
|
ErrorLatched bool
|
||||||
|
LastError *string
|
||||||
|
Charge string
|
||||||
|
CleanType string
|
||||||
|
Paused bool
|
||||||
|
Fan string
|
||||||
|
}
|
||||||
|
|
||||||
|
// StateDocument is the HA state payload, field order state then fan_speed.
|
||||||
|
type StateDocument struct {
|
||||||
|
State string `json:"state"`
|
||||||
|
FanSpeed string `json:"fan_speed"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// LifespanTotal preserves the reported totals; their unit is unknown.
|
||||||
|
type LifespanTotal struct {
|
||||||
|
SideBrush *int `json:"side_brush"`
|
||||||
|
MainBrush *int `json:"main_brush"`
|
||||||
|
Filter *int `json:"filter"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ScheduleAction is the inner ctl of a schedule entry.
|
||||||
|
type ScheduleAction struct {
|
||||||
|
TD string `json:"td"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Schedule is one §10.3 schedule entry. Phase 05 is its only writer.
|
||||||
|
type Schedule struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
On bool `json:"on"`
|
||||||
|
Time string `json:"time"`
|
||||||
|
Repeat string `json:"repeat"`
|
||||||
|
Flag string `json:"flag"`
|
||||||
|
Action ScheduleAction `json:"action"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// AttributeDocument is the full §10.3 json_attributes object. Pointer fields
|
||||||
|
// publish JSON null until a report fills them.
|
||||||
|
type AttributeDocument struct {
|
||||||
|
BatteryLevel *int `json:"battery_level"`
|
||||||
|
SideBrush *int `json:"side_brush"`
|
||||||
|
MainBrush *int `json:"main_brush"`
|
||||||
|
Filter *int `json:"filter"`
|
||||||
|
LifespanTotal LifespanTotal `json:"lifespan_total"`
|
||||||
|
CleanType *string `json:"clean_type"`
|
||||||
|
ChargeState *string `json:"charge_state"`
|
||||||
|
LastError *string `json:"last_error"`
|
||||||
|
LastCommandError *string `json:"last_command_error"`
|
||||||
|
Schedules []Schedule `json:"schedules"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Snapshot is the only retained view of a robot. The actor is its sole
|
||||||
|
// mutator; RepublishFunc receives a Clone so slices and pointers cannot race.
|
||||||
|
type Snapshot struct {
|
||||||
|
Facts Facts `json:"-"`
|
||||||
|
State StateDocument `json:"-"`
|
||||||
|
Attributes AttributeDocument `json:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RepublishFunc is invoked after every successful Apply with both documents
|
||||||
|
// filled. Phase 04 assigns the publishing function; nil skips IO.
|
||||||
|
type RepublishFunc func(ctx context.Context, snap Snapshot)
|
||||||
|
|
||||||
|
// ScheduleParser parses the ctl inner XML of a GetSched result or a Sched2
|
||||||
|
// push into the complete schedules array. Phase 05 assigns the hook.
|
||||||
|
type ScheduleParser func(inner []byte) ([]Schedule, error)
|
||||||
|
|
||||||
|
// LifespanData holds the parsed result of a GetLifeSpan query.
|
||||||
|
type LifespanData struct {
|
||||||
|
Target string // e.g. "side_brush", "main_brush", "filter"
|
||||||
|
Val int
|
||||||
|
Total int
|
||||||
|
}
|
||||||
|
|
||||||
|
// LifespanParser consumes the ctl attributes of a GetLifeSpan result.
|
||||||
|
// Phase 05 assigns the hook.
|
||||||
|
type LifespanParser func(td string, attrs map[string]string) (LifespanData, error)
|
||||||
|
|
||||||
|
// ParseSchedules and ParseLifespan are the phase 05 hooks. Nil means the
|
||||||
|
// corresponding result still completes its cid and leaves the fields empty.
|
||||||
|
var (
|
||||||
|
ParseSchedules ScheduleParser
|
||||||
|
ParseLifespan LifespanParser
|
||||||
|
)
|
||||||
|
|
||||||
|
// ValidationError marks a reported value outside its allowed range. Apply
|
||||||
|
// leaves the offending field untouched and returns this error so the actor
|
||||||
|
// can emit an unparsed diagnostic.
|
||||||
|
type ValidationError struct {
|
||||||
|
Field string
|
||||||
|
Value string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *ValidationError) Error() string {
|
||||||
|
return "invalid " + e.Field + " " + strconv.Quote(e.Value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSnapshot returns the initial snapshot: fan standard, HA state idle,
|
||||||
|
// schedules an empty array.
|
||||||
|
func NewSnapshot() Snapshot {
|
||||||
|
s := Snapshot{
|
||||||
|
Facts: Facts{Fan: "standard"},
|
||||||
|
Attributes: AttributeDocument{
|
||||||
|
Schedules: []Schedule{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s.rebuild()
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// rebuild rewrites the derived fields of both documents. Stored fields such
|
||||||
|
// as battery_level, consumables, last_command_error and schedules persist.
|
||||||
|
func (s *Snapshot) rebuild() {
|
||||||
|
s.State.State = Derive(s.Facts)
|
||||||
|
s.State.FanSpeed = s.Facts.Fan
|
||||||
|
s.Attributes.CleanType = stringOrNil(s.Facts.CleanType)
|
||||||
|
s.Attributes.ChargeState = stringOrNil(s.Facts.Charge)
|
||||||
|
s.Attributes.LastError = s.Facts.LastError
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply folds one inbound payload into the snapshot and rebuilds both
|
||||||
|
// documents. An out-of-range battery or an unknown fan speed leaves that
|
||||||
|
// value untouched and returns a ValidationError. Result stanzas carry their
|
||||||
|
// command name in in.TD so a SetCleanSpeed result without a speed echo can
|
||||||
|
// still store requestedFan.
|
||||||
|
func Apply(snap *Snapshot, in ctl.Inbound, requestedFan string) error {
|
||||||
|
var firstErr error
|
||||||
|
fail := func(err error) {
|
||||||
|
if firstErr == nil {
|
||||||
|
firstErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if in.Kind != ctl.KindBattery && !knownTD(in.TD) {
|
||||||
|
return &ValidationError{Field: "td", Value: in.TD}
|
||||||
|
}
|
||||||
|
|
||||||
|
if in.TD == "error" {
|
||||||
|
if in.Errno != nil && *in.Errno == "100" {
|
||||||
|
snap.Facts.ErrorLatched = false
|
||||||
|
snap.Facts.LastError = nil
|
||||||
|
} else {
|
||||||
|
snap.Facts.ErrorLatched = true
|
||||||
|
snap.Facts.LastError = in.Errno
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if in.TD == "Idle" {
|
||||||
|
snap.Facts.Charge = "Idle"
|
||||||
|
}
|
||||||
|
|
||||||
|
if in.CleanAttrs != nil {
|
||||||
|
if v := in.CleanAttrs["type"]; v != "" {
|
||||||
|
snap.Facts.CleanType = v
|
||||||
|
}
|
||||||
|
if v := in.CleanAttrs["speed"]; v != "" {
|
||||||
|
if validFan(v) {
|
||||||
|
snap.Facts.Fan = v
|
||||||
|
} else {
|
||||||
|
fail(&ValidationError{Field: "clean speed", Value: v})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var keys []string
|
||||||
|
switch in.TD {
|
||||||
|
case "CleanReport":
|
||||||
|
keys = []string{"st", "rsn"}
|
||||||
|
case "GetCleanState":
|
||||||
|
keys = []string{"t", "a"}
|
||||||
|
}
|
||||||
|
var nonblank []string
|
||||||
|
for _, k := range keys {
|
||||||
|
if v := in.CleanAttrs[k]; strings.TrimSpace(v) != "" {
|
||||||
|
nonblank = append(nonblank, fmt.Sprintf("%s=%q", k, v))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(nonblank) > 0 {
|
||||||
|
fail(fmt.Errorf("nonblank clean %s", strings.Join(nonblank, " ")))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if in.ChargeAttrs != nil {
|
||||||
|
if v := in.ChargeAttrs["type"]; v != "" {
|
||||||
|
snap.Facts.Charge = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if in.BatteryPower != "" {
|
||||||
|
if n, err := strconv.Atoi(in.BatteryPower); err == nil && n >= 0 && n <= 100 {
|
||||||
|
snap.Attributes.BatteryLevel = &n
|
||||||
|
} else {
|
||||||
|
fail(&ValidationError{Field: "battery power", Value: in.BatteryPower})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if in.TD == "GetCleanSpeed" {
|
||||||
|
if v := in.Attrs["speed"]; v != "" {
|
||||||
|
if validFan(v) {
|
||||||
|
snap.Facts.Fan = v
|
||||||
|
} else {
|
||||||
|
fail(&ValidationError{Field: "clean speed", Value: v})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if in.TD == "SetCleanSpeed" && requestedFan != "" {
|
||||||
|
if validFan(requestedFan) {
|
||||||
|
snap.Facts.Fan = requestedFan
|
||||||
|
} else {
|
||||||
|
fail(&ValidationError{Field: "clean speed", Value: requestedFan})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (in.TD == "GetSched" || in.TD == "Sched2") && ParseSchedules != nil {
|
||||||
|
schedules, err := ParseSchedules(in.Inner)
|
||||||
|
if schedules == nil {
|
||||||
|
snap.Attributes.Schedules = []Schedule{}
|
||||||
|
} else {
|
||||||
|
snap.Attributes.Schedules = schedules
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
fail(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if in.TD == "GetLifeSpan" && ParseLifespan != nil {
|
||||||
|
data, err := ParseLifespan(in.TD, in.Attrs)
|
||||||
|
if err != nil {
|
||||||
|
fail(err)
|
||||||
|
} else {
|
||||||
|
v := data.Val
|
||||||
|
t := data.Total
|
||||||
|
switch data.Target {
|
||||||
|
case "side_brush":
|
||||||
|
snap.Attributes.SideBrush = &v
|
||||||
|
snap.Attributes.LifespanTotal.SideBrush = &t
|
||||||
|
case "main_brush":
|
||||||
|
snap.Attributes.MainBrush = &v
|
||||||
|
snap.Attributes.LifespanTotal.MainBrush = &t
|
||||||
|
case "filter":
|
||||||
|
snap.Attributes.Filter = &v
|
||||||
|
snap.Attributes.LifespanTotal.Filter = &t
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
snap.rebuild()
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCommandError sets last_command_error and rebuilds both documents. It
|
||||||
|
// never touches last_error.
|
||||||
|
func (s *Snapshot) SetCommandError(value *string) {
|
||||||
|
s.Attributes.LastCommandError = value
|
||||||
|
s.rebuild()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clone returns a deep-enough copy that no slice or pointer is shared with
|
||||||
|
// the original, so a published snapshot cannot race later mutation.
|
||||||
|
func (s Snapshot) Clone() Snapshot {
|
||||||
|
c := s
|
||||||
|
c.Facts.LastError = cloneStr(s.Facts.LastError)
|
||||||
|
a := &c.Attributes
|
||||||
|
a.BatteryLevel = cloneInt(a.BatteryLevel)
|
||||||
|
a.SideBrush = cloneInt(a.SideBrush)
|
||||||
|
a.MainBrush = cloneInt(a.MainBrush)
|
||||||
|
a.Filter = cloneInt(a.Filter)
|
||||||
|
a.LifespanTotal.SideBrush = cloneInt(a.LifespanTotal.SideBrush)
|
||||||
|
a.LifespanTotal.MainBrush = cloneInt(a.LifespanTotal.MainBrush)
|
||||||
|
a.LifespanTotal.Filter = cloneInt(a.LifespanTotal.Filter)
|
||||||
|
a.CleanType = cloneStr(a.CleanType)
|
||||||
|
a.ChargeState = cloneStr(a.ChargeState)
|
||||||
|
a.LastError = cloneStr(a.LastError)
|
||||||
|
a.LastCommandError = cloneStr(a.LastCommandError)
|
||||||
|
if a.Schedules != nil {
|
||||||
|
a.Schedules = append([]Schedule{}, a.Schedules...)
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func validFan(v string) bool { return v == "standard" || v == "strong" }
|
||||||
|
|
||||||
|
// knownTD reports whether td is a value this phase understands: every
|
||||||
|
// outgoing command name and every supported push.
|
||||||
|
func knownTD(td string) bool {
|
||||||
|
switch td {
|
||||||
|
case "SetTime", "GetBatteryInfo", "GetCleanState", "GetChargeState",
|
||||||
|
"GetCleanSpeed", "GetSched", "GetLifeSpan",
|
||||||
|
"Clean", "Charge", "PlaySound", "SetCleanSpeed", "Move",
|
||||||
|
"Sched2", "CleanReport", "ChargeState", "BatteryInfo", "error", "Idle",
|
||||||
|
"AddSched", "ModSched", "DelSched":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func stringOrNil(v string) *string {
|
||||||
|
if v == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
vv := v
|
||||||
|
return &vv
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneStr(p *string) *string {
|
||||||
|
if p == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
v := *p
|
||||||
|
return &v
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneInt(p *int) *int {
|
||||||
|
if p == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
v := *p
|
||||||
|
return &v
|
||||||
|
}
|
||||||
@@ -0,0 +1,223 @@
|
|||||||
|
// Package session owns the XMPP generation registry, observer bus, and the
|
||||||
|
// diagnostic sink type used across the XMPP and robot layers.
|
||||||
|
package session
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrStale is returned by Registry.Send when the requested generation is no
|
||||||
|
// longer current for the JID.
|
||||||
|
var ErrStale = errors.New("stale generation")
|
||||||
|
|
||||||
|
// Reasons carried by DownEvent.
|
||||||
|
const (
|
||||||
|
ReasonTCPClose = "tcp-close"
|
||||||
|
ReasonStreamClose = "stream-close"
|
||||||
|
ReasonPingTimeout = "ping-timeout"
|
||||||
|
ReasonMalformedXML = "malformed-xml"
|
||||||
|
ReasonReplaced = "replaced"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Diagnostic directions.
|
||||||
|
const (
|
||||||
|
DirectionIn = "in"
|
||||||
|
DirectionOut = "out"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Registry holds at most one current generation per full bot JID. Bind
|
||||||
|
// atomically replaces any existing slot; the old slot is closed and the
|
||||||
|
// previous generation becomes stale.
|
||||||
|
type Registry struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
current map[string]*slot
|
||||||
|
}
|
||||||
|
|
||||||
|
type slot struct {
|
||||||
|
gen uint64
|
||||||
|
send func([]byte) error
|
||||||
|
close func()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRegistry creates an empty registry.
|
||||||
|
func NewRegistry() *Registry {
|
||||||
|
return &Registry{current: make(map[string]*slot)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bind installs a new generation for jid. It returns the new generation and
|
||||||
|
// whether an existing slot was replaced. The old slot is closed after the
|
||||||
|
// registry has been updated.
|
||||||
|
func (r *Registry) Bind(jid string, send func([]byte) error, close func()) (gen uint64, replaced bool) {
|
||||||
|
r.mu.Lock()
|
||||||
|
old := r.current[jid]
|
||||||
|
var next uint64
|
||||||
|
if old != nil {
|
||||||
|
replaced = true
|
||||||
|
next = old.gen + 1
|
||||||
|
} else {
|
||||||
|
next = 1
|
||||||
|
}
|
||||||
|
r.current[jid] = &slot{gen: next, send: send, close: close}
|
||||||
|
r.mu.Unlock()
|
||||||
|
|
||||||
|
if old != nil && old.close != nil {
|
||||||
|
old.close()
|
||||||
|
}
|
||||||
|
return next, replaced
|
||||||
|
}
|
||||||
|
|
||||||
|
// Current returns the current generation for jid, if any.
|
||||||
|
func (r *Registry) Current(jid string) (gen uint64, ok bool) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
s := r.current[jid]
|
||||||
|
if s == nil {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
return s.gen, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send delivers xml to the current generation. It returns ErrStale if the
|
||||||
|
// generation is no longer current. The lock is released before calling the
|
||||||
|
// slot's send function so a slow or blocked write cannot stall the registry.
|
||||||
|
func (r *Registry) Send(jid string, gen uint64, xml []byte) error {
|
||||||
|
r.mu.Lock()
|
||||||
|
s := r.current[jid]
|
||||||
|
r.mu.Unlock()
|
||||||
|
if s == nil || s.gen != gen {
|
||||||
|
return ErrStale
|
||||||
|
}
|
||||||
|
return s.send(xml)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove deletes the slot for jid only if it matches gen. It returns true when
|
||||||
|
// the slot existed and matched, which a read loop uses after emitting its own
|
||||||
|
// SessionDown for the current generation.
|
||||||
|
func (r *Registry) Remove(jid string, gen uint64) bool {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
s := r.current[jid]
|
||||||
|
if s != nil && s.gen == gen {
|
||||||
|
delete(r.current, jid)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// CloseAll calls every current slot's close function without removing it. The
|
||||||
|
// read loops observe the closed socket and emit SessionDown. New slots cannot
|
||||||
|
// be bound while CloseAll runs because the caller owns shutdown order.
|
||||||
|
func (r *Registry) CloseAll() []func() {
|
||||||
|
r.mu.Lock()
|
||||||
|
closers := make([]func(), 0, len(r.current))
|
||||||
|
for _, s := range r.current {
|
||||||
|
closers = append(closers, s.close)
|
||||||
|
}
|
||||||
|
r.mu.Unlock()
|
||||||
|
for _, c := range closers {
|
||||||
|
c()
|
||||||
|
}
|
||||||
|
return closers
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadyEvent is emitted when a full handshake reaches READY.
|
||||||
|
type ReadyEvent struct {
|
||||||
|
Generation uint64
|
||||||
|
JID string
|
||||||
|
Serial string
|
||||||
|
Send func([]byte) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// DownEvent is emitted when a generation ends.
|
||||||
|
type DownEvent struct {
|
||||||
|
Generation uint64
|
||||||
|
JID string
|
||||||
|
Serial string
|
||||||
|
Reason string
|
||||||
|
}
|
||||||
|
|
||||||
|
// StanzaEvent is emitted for every non-ping post-READY stanza.
|
||||||
|
type StanzaEvent struct {
|
||||||
|
Generation uint64
|
||||||
|
JID string
|
||||||
|
Serial string
|
||||||
|
Stanza []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// Observer receives session lifecycle events. Implementations may implement
|
||||||
|
// only the methods they care about by providing no-op bodies.
|
||||||
|
type Observer interface {
|
||||||
|
SessionReady(e ReadyEvent)
|
||||||
|
AnnounceOK(e ReadyEvent)
|
||||||
|
SessionDown(e DownEvent)
|
||||||
|
Stanza(e StanzaEvent)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bus fans observer calls out to registered observers.
|
||||||
|
type Bus struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
obs []Observer
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBus creates an empty bus.
|
||||||
|
func NewBus() *Bus {
|
||||||
|
return &Bus{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register adds an observer. Observers are called in registration order.
|
||||||
|
func (b *Bus) Register(o Observer) {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
b.obs = append(b.obs, o)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bus) SessionReady(e ReadyEvent) {
|
||||||
|
b.mu.RLock()
|
||||||
|
obs := append([]Observer(nil), b.obs...)
|
||||||
|
b.mu.RUnlock()
|
||||||
|
for _, o := range obs {
|
||||||
|
o.SessionReady(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bus) AnnounceOK(e ReadyEvent) {
|
||||||
|
b.mu.RLock()
|
||||||
|
obs := append([]Observer(nil), b.obs...)
|
||||||
|
b.mu.RUnlock()
|
||||||
|
for _, o := range obs {
|
||||||
|
o.AnnounceOK(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bus) SessionDown(e DownEvent) {
|
||||||
|
b.mu.RLock()
|
||||||
|
obs := append([]Observer(nil), b.obs...)
|
||||||
|
b.mu.RUnlock()
|
||||||
|
for _, o := range obs {
|
||||||
|
o.SessionDown(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Bus) Stanza(e StanzaEvent) {
|
||||||
|
b.mu.RLock()
|
||||||
|
obs := append([]Observer(nil), b.obs...)
|
||||||
|
b.mu.RUnlock()
|
||||||
|
for _, o := range obs {
|
||||||
|
o.Stanza(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Diagnostic carries a redacted XMPP fragment for the raw diagnostic topic.
|
||||||
|
// Auth material never reaches a Diagnostic.
|
||||||
|
type Diagnostic struct {
|
||||||
|
Serial string
|
||||||
|
Direction string
|
||||||
|
Kind string
|
||||||
|
Generation uint64
|
||||||
|
XML []byte
|
||||||
|
Reason string
|
||||||
|
}
|
||||||
|
|
||||||
|
// DiagnosticSink receives diagnostics. A nil sink drops the event.
|
||||||
|
type DiagnosticSink func(Diagnostic)
|
||||||
@@ -0,0 +1,128 @@
|
|||||||
|
package session
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBindGenerations(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
gen1, replaced1 := r.Bind("a@b/c", nil, nil)
|
||||||
|
if gen1 != 1 || replaced1 {
|
||||||
|
t.Fatalf("first bind: gen=%d replaced=%v, want 1 false", gen1, replaced1)
|
||||||
|
}
|
||||||
|
gen2, replaced2 := r.Bind("a@b/c", nil, nil)
|
||||||
|
if gen2 != 2 || !replaced2 {
|
||||||
|
t.Fatalf("second bind: gen=%d replaced=%v, want 2 true", gen2, replaced2)
|
||||||
|
}
|
||||||
|
g, ok := r.Current("a@b/c")
|
||||||
|
if !ok || g != 2 {
|
||||||
|
t.Fatalf("Current = %d %v, want 2 true", g, ok)
|
||||||
|
}
|
||||||
|
_, ok = r.Current("other@b/c")
|
||||||
|
if ok {
|
||||||
|
t.Fatalf("unknown JID should not be current")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReplaceClosesOldSlot(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
var closedMu sync.Mutex
|
||||||
|
var closedGen uint64
|
||||||
|
closeFn := func(gen uint64) func() {
|
||||||
|
return func() {
|
||||||
|
closedMu.Lock()
|
||||||
|
closedGen = gen
|
||||||
|
closedMu.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
r.Bind("a@b/c", func([]byte) error { return nil }, closeFn(1))
|
||||||
|
_, replaced := r.Bind("a@b/c", func([]byte) error { return nil }, closeFn(2))
|
||||||
|
if !replaced {
|
||||||
|
t.Fatalf("expected replaced")
|
||||||
|
}
|
||||||
|
|
||||||
|
closedMu.Lock()
|
||||||
|
g := closedGen
|
||||||
|
closedMu.Unlock()
|
||||||
|
if g != 1 {
|
||||||
|
t.Fatalf("old slot closed gen=%d, want 1", g)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendStale(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
var sent []byte
|
||||||
|
r.Bind("a@b/c", func(b []byte) error { sent = append([]byte(nil), b...); return nil }, func() {})
|
||||||
|
|
||||||
|
err := r.Send("a@b/c", 1, []byte("hello"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Send current gen: %v", err)
|
||||||
|
}
|
||||||
|
if string(sent) != "hello" {
|
||||||
|
t.Fatalf("sent = %q", sent)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.Send("a@b/c", 2, []byte("x")); !errors.Is(err, ErrStale) {
|
||||||
|
t.Fatalf("stale gen error = %v, want ErrStale", err)
|
||||||
|
}
|
||||||
|
if err := r.Send("unknown@b/c", 1, []byte("x")); !errors.Is(err, ErrStale) {
|
||||||
|
t.Fatalf("unknown jid error = %v, want ErrStale", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBusFanOut(t *testing.T) {
|
||||||
|
b := NewBus()
|
||||||
|
var ready, announce, down, stanza int
|
||||||
|
o1 := &testObserver{
|
||||||
|
onReady: func() { ready++ },
|
||||||
|
onAnnounce: func() { announce++ },
|
||||||
|
onDown: func() { down++ },
|
||||||
|
onStanza: func() { stanza++ },
|
||||||
|
}
|
||||||
|
o2 := &testObserver{
|
||||||
|
onReady: func() { ready++ },
|
||||||
|
onAnnounce: func() { announce++ },
|
||||||
|
onDown: func() { down++ },
|
||||||
|
onStanza: func() { stanza++ },
|
||||||
|
}
|
||||||
|
b.Register(o1)
|
||||||
|
b.Register(o2)
|
||||||
|
|
||||||
|
b.SessionReady(ReadyEvent{})
|
||||||
|
b.AnnounceOK(ReadyEvent{})
|
||||||
|
b.SessionDown(DownEvent{})
|
||||||
|
b.Stanza(StanzaEvent{})
|
||||||
|
|
||||||
|
if ready != 2 || announce != 2 || down != 2 || stanza != 2 {
|
||||||
|
t.Fatalf("counts = %d/%d/%d/%d, want all 2", ready, announce, down, stanza)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveCurrentSlot(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
r.Bind("a@b/c", nil, func() {})
|
||||||
|
if !r.Remove("a@b/c", 1) {
|
||||||
|
t.Fatalf("Remove current gen should succeed")
|
||||||
|
}
|
||||||
|
if _, ok := r.Current("a@b/c"); ok {
|
||||||
|
t.Fatalf("JID should no longer be current")
|
||||||
|
}
|
||||||
|
if r.Remove("a@b/c", 1) {
|
||||||
|
t.Fatalf("Remove same gen again should fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type testObserver struct {
|
||||||
|
onReady func()
|
||||||
|
onAnnounce func()
|
||||||
|
onDown func()
|
||||||
|
onStanza func()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *testObserver) SessionReady(_ ReadyEvent) { o.onReady() }
|
||||||
|
func (o *testObserver) AnnounceOK(_ ReadyEvent) { o.onAnnounce() }
|
||||||
|
func (o *testObserver) SessionDown(_ DownEvent) { o.onDown() }
|
||||||
|
func (o *testObserver) Stanza(_ StanzaEvent) { o.onStanza() }
|
||||||
@@ -0,0 +1,509 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/xml"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
nsClient = "jabber:client"
|
||||||
|
nsStream = "http://etherx.jabber.org/streams"
|
||||||
|
nsSASL = "urn:ietf:params:xml:ns:xmpp-sasl"
|
||||||
|
nsBind = "urn:ietf:params:xml:ns:xmpp-bind"
|
||||||
|
nsSession = "urn:ietf:params:xml:ns:xmpp-session"
|
||||||
|
nsTLS = "urn:ietf:params:xml:ns:xmpp-tls"
|
||||||
|
nsIQAuth = "http://jabber.org/features/iq-auth"
|
||||||
|
nsPing = "urn:xmpp:ping"
|
||||||
|
)
|
||||||
|
|
||||||
|
// handshake states.
|
||||||
|
const (
|
||||||
|
hsNeedAuth = iota
|
||||||
|
hsNeedBind
|
||||||
|
hsNeedSession
|
||||||
|
hsNeedHelloWorld
|
||||||
|
hsReady
|
||||||
|
)
|
||||||
|
|
||||||
|
// Conn handles one plaintext XMPP connection from accept to close.
|
||||||
|
type Conn struct {
|
||||||
|
netConn net.Conn
|
||||||
|
cfg config.Config
|
||||||
|
registry *session.Registry
|
||||||
|
bus *session.Bus
|
||||||
|
clock Clock
|
||||||
|
diag session.DiagnosticSink
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
closed bool
|
||||||
|
|
||||||
|
tok *Tokenizer
|
||||||
|
hsState int
|
||||||
|
streamID string
|
||||||
|
domain string
|
||||||
|
authcid string
|
||||||
|
resource string
|
||||||
|
fullJID string
|
||||||
|
serial string
|
||||||
|
gen uint64
|
||||||
|
|
||||||
|
pingLoop *pingLoop
|
||||||
|
}
|
||||||
|
|
||||||
|
// serve runs the connection until it ends. It owns the read loop and cleanup.
|
||||||
|
func (c *Conn) serve() {
|
||||||
|
c.tok = NewTokenizer()
|
||||||
|
defer c.close()
|
||||||
|
buf := make([]byte, 4096)
|
||||||
|
for {
|
||||||
|
n, err := c.netConn.Read(buf)
|
||||||
|
if n > 0 {
|
||||||
|
evs := c.tok.Feed(buf[:n])
|
||||||
|
for _, ev := range evs {
|
||||||
|
if !c.handleEvent(ev) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
if err != io.EOF && !isClosed(err) {
|
||||||
|
slog.Debug("xmpp read error", "jid", c.fullJID, "err", err)
|
||||||
|
}
|
||||||
|
c.endSession(session.ReasonTCPClose)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isClosed(err error) bool {
|
||||||
|
if err == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if netErr, ok := err.(net.Error); ok {
|
||||||
|
return netErr.Timeout() || strings.Contains(err.Error(), "use of closed network connection")
|
||||||
|
}
|
||||||
|
return strings.Contains(err.Error(), "use of closed network connection")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleEvent(ev Event) bool {
|
||||||
|
switch ev.Kind {
|
||||||
|
case StreamOpen:
|
||||||
|
return c.handleStreamOpen(ev.Attrs)
|
||||||
|
case Stanza:
|
||||||
|
if !c.isAuthStanza(ev.Data) {
|
||||||
|
logStanza("in", ev.Data)
|
||||||
|
c.emitDiag(session.DirectionIn, "stanza", ev.Data, "")
|
||||||
|
}
|
||||||
|
if c.hsState == hsReady {
|
||||||
|
return c.handlePostReadyStanza(ev.Data)
|
||||||
|
}
|
||||||
|
return c.handlePreReadyStanza(ev.Data)
|
||||||
|
case StreamClose:
|
||||||
|
c.endSession(session.ReasonStreamClose)
|
||||||
|
return false
|
||||||
|
case TokenError:
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) isAuthStanza(b []byte) bool {
|
||||||
|
space, local, err := rootSpaceLocal(b)
|
||||||
|
return err == nil && space == nsSASL && local == "auth"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleStreamOpen(attrs map[string]string) bool {
|
||||||
|
domain := attrs["to"]
|
||||||
|
if domain == "" {
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
c.domain = domain
|
||||||
|
if c.streamID == "" {
|
||||||
|
c.streamID = newStreamID()
|
||||||
|
}
|
||||||
|
|
||||||
|
var payload string
|
||||||
|
switch c.hsState {
|
||||||
|
case hsNeedAuth:
|
||||||
|
payload = streamOpen(c.streamID, c.domain) +
|
||||||
|
`<stream:features>` +
|
||||||
|
`<auth xmlns="http://jabber.org/features/iq-auth"/>` +
|
||||||
|
`<starttls xmlns="urn:ietf:params:xml:ns:xmpp-tls"><required/></starttls>` +
|
||||||
|
`<mechanisms xmlns="urn:ietf:params:xml:ns:xmpp-sasl"><mechanism>PLAIN</mechanism></mechanisms>` +
|
||||||
|
`</stream:features>`
|
||||||
|
case hsNeedBind:
|
||||||
|
payload = streamOpen(c.streamID, c.domain) +
|
||||||
|
`<stream:features>` +
|
||||||
|
`<bind xmlns="urn:ietf:params:xml:ns:xmpp-bind"/>` +
|
||||||
|
`<session xmlns="urn:ietf:params:xml:ns:xmpp-session"/>` +
|
||||||
|
`</stream:features>`
|
||||||
|
default:
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if err := c.writeRaw([]byte(payload)); err != nil {
|
||||||
|
c.endSession(session.ReasonTCPClose)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func streamOpen(id, domain string) string {
|
||||||
|
return fmt.Sprintf(`<stream:stream xmlns:stream="http://etherx.jabber.org/streams" xmlns="jabber:client" version="1.0" id="%s" from="%s">`, id, domain)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handlePreReadyStanza(stanza []byte) bool {
|
||||||
|
space, local, err := rootSpaceLocal(stanza)
|
||||||
|
if err != nil {
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch local {
|
||||||
|
case "starttls":
|
||||||
|
// STARTTLS is advertised but the robot ignores it; keep reading.
|
||||||
|
return true
|
||||||
|
case "auth":
|
||||||
|
switch space {
|
||||||
|
case nsIQAuth:
|
||||||
|
// iq-auth is advertised but not supported. Fail without logging payload.
|
||||||
|
slog.Info("sasl invalid mechanism")
|
||||||
|
_ = c.writeStanza([]byte(`<failure xmlns="http://jabber.org/features/iq-auth"><not-authorized/></failure>`))
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
case nsSASL:
|
||||||
|
return c.handleAuth(stanza)
|
||||||
|
}
|
||||||
|
case "iq":
|
||||||
|
return c.handleIQStanza(stanza)
|
||||||
|
case "presence":
|
||||||
|
return c.handleHelloWorld(stanza)
|
||||||
|
}
|
||||||
|
// Unknown well-formed pre-READY stanza: log and continue.
|
||||||
|
logStanza("in", stanza)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleAuth(stanza []byte) bool {
|
||||||
|
var auth saslAuth
|
||||||
|
if err := xml.Unmarshal(stanza, &auth); err != nil {
|
||||||
|
slog.Info("sasl malformed")
|
||||||
|
_ = c.writeStanza(SASLFailureXML("malformed-request"))
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if auth.Mechanism != "PLAIN" {
|
||||||
|
slog.Info("sasl invalid mechanism")
|
||||||
|
_ = c.writeStanza(SASLFailureXML("invalid-mechanism"))
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
parsed, err := ParseSASLPlain(auth.Chardata)
|
||||||
|
if err != nil {
|
||||||
|
slog.Info("sasl malformed")
|
||||||
|
_ = c.writeStanza(SASLFailureXML("malformed-request"))
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
c.authcid = parsed.Authcid
|
||||||
|
slog.Info("sasl authenticated", "authcid", c.authcid)
|
||||||
|
_ = c.writeStanza(SASLSuccessXML)
|
||||||
|
c.tok.ExpectNewStream()
|
||||||
|
c.hsState = hsNeedBind
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleIQStanza(stanza []byte) bool {
|
||||||
|
// Try bind first.
|
||||||
|
var bindReq iqBind
|
||||||
|
if err := xml.Unmarshal(stanza, &bindReq); err == nil && bindReq.Bind.XMLName.Local == "bind" {
|
||||||
|
return c.handleBind(bindReq)
|
||||||
|
}
|
||||||
|
// Then session.
|
||||||
|
var sessReq iqSession
|
||||||
|
if err := xml.Unmarshal(stanza, &sessReq); err == nil && sessReq.Session.XMLName.Local == "session" {
|
||||||
|
return c.handleSession(sessReq)
|
||||||
|
}
|
||||||
|
logStanza("in", stanza)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleBind(req iqBind) bool {
|
||||||
|
if req.Type != "set" {
|
||||||
|
logStanza("in", []byte(fmt.Sprintf("bind iq type=%s", req.Type)))
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
c.resource = req.Bind.Resource
|
||||||
|
if c.resource == "" {
|
||||||
|
c.resource = "atom"
|
||||||
|
}
|
||||||
|
c.fullJID = fmt.Sprintf("%s@%s/%s", c.authcid, c.domain, c.resource)
|
||||||
|
c.serial = c.authcid
|
||||||
|
|
||||||
|
sendFn := func(b []byte) error { return c.writeRaw(b) }
|
||||||
|
closeFn := func() { c.closeWith(nil) }
|
||||||
|
gen, replaced := c.registry.Bind(c.fullJID, sendFn, closeFn)
|
||||||
|
c.gen = gen
|
||||||
|
|
||||||
|
if replaced {
|
||||||
|
c.bus.SessionDown(session.DownEvent{
|
||||||
|
Generation: gen - 1,
|
||||||
|
JID: c.fullJID,
|
||||||
|
Serial: c.serial,
|
||||||
|
Reason: session.ReasonReplaced,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
result := fmt.Sprintf(`<iq type="result" id="%s"><bind xmlns="urn:ietf:params:xml:ns:xmpp-bind"><jid>%s</jid></bind></iq>`,
|
||||||
|
escapeXMLAttr(req.ID), c.fullJID)
|
||||||
|
if err := c.writeStanza([]byte(result)); err != nil {
|
||||||
|
c.endSession(session.ReasonTCPClose)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
c.hsState = hsNeedSession
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleSession(req iqSession) bool {
|
||||||
|
if req.Type != "set" {
|
||||||
|
logStanza("in", []byte(fmt.Sprintf("session iq type=%s", req.Type)))
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
result := fmt.Sprintf(`<iq type="result" id="%s"/>`, escapeXMLAttr(req.ID))
|
||||||
|
if err := c.writeStanza([]byte(result)); err != nil {
|
||||||
|
c.endSession(session.ReasonTCPClose)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
c.hsState = hsNeedHelloWorld
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handleHelloWorld(stanza []byte) bool {
|
||||||
|
var p presenceHello
|
||||||
|
if err := xml.Unmarshal(stanza, &p); err != nil {
|
||||||
|
c.endSession(session.ReasonMalformedXML)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(p.Status) != "hello world" {
|
||||||
|
// Not the READY presence; keep waiting.
|
||||||
|
logStanza("in", stanza)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
dummy := fmt.Sprintf(`<presence to="%s"> dummy </presence>`, escapeXMLAttr(c.fullJID))
|
||||||
|
if err := c.writeStanza([]byte(dummy)); err != nil {
|
||||||
|
c.endSession(session.ReasonTCPClose)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
c.hsState = hsReady
|
||||||
|
|
||||||
|
c.bus.SessionReady(session.ReadyEvent{
|
||||||
|
Generation: c.gen,
|
||||||
|
JID: c.fullJID,
|
||||||
|
Serial: c.serial,
|
||||||
|
Send: func(b []byte) error { return c.registry.Send(c.fullJID, c.gen, b) },
|
||||||
|
})
|
||||||
|
|
||||||
|
c.pingLoop = newPingLoop(c)
|
||||||
|
go c.pingLoop.run()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) handlePostReadyStanza(stanza []byte) bool {
|
||||||
|
if !c.isCurrent() {
|
||||||
|
// Generation was replaced; stop processing buffered data.
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bot domain ping.
|
||||||
|
var ping iqPing
|
||||||
|
if err := xml.Unmarshal(stanza, &ping); err == nil && ping.Type == "get" && ping.Ping.XMLName.Space == nsPing {
|
||||||
|
reply := fmt.Sprintf(`<iq type="result" to="%s" from="%s" id="%s"/>`,
|
||||||
|
escapeXMLAttr(ping.From), escapeXMLAttr(ping.To), escapeXMLAttr(ping.ID))
|
||||||
|
_ = c.writeStanza([]byte(reply))
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Controller ping result.
|
||||||
|
var result iqResult
|
||||||
|
if err := xml.Unmarshal(stanza, &result); err == nil && result.Type == "result" && result.ID != "" {
|
||||||
|
if c.pingLoop != nil && c.pingLoop.result(result.ID) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Any other stanza is forwarded to observers.
|
||||||
|
c.bus.Stanza(session.StanzaEvent{
|
||||||
|
Generation: c.gen,
|
||||||
|
JID: c.fullJID,
|
||||||
|
Serial: c.serial,
|
||||||
|
Stanza: append([]byte(nil), stanza...),
|
||||||
|
})
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) isCurrent() bool {
|
||||||
|
gen, ok := c.registry.Current(c.fullJID)
|
||||||
|
return ok && gen == c.gen
|
||||||
|
}
|
||||||
|
|
||||||
|
// endSession closes the connection and, if this generation is still current,
|
||||||
|
// emits SessionDown and removes the registry slot.
|
||||||
|
func (c *Conn) endSession(reason string) {
|
||||||
|
c.close()
|
||||||
|
if c.gen == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
currentGen, ok := c.registry.Current(c.fullJID)
|
||||||
|
if !ok || currentGen != c.gen {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.bus.SessionDown(session.DownEvent{
|
||||||
|
Generation: c.gen,
|
||||||
|
JID: c.fullJID,
|
||||||
|
Serial: c.serial,
|
||||||
|
Reason: reason,
|
||||||
|
})
|
||||||
|
c.registry.Remove(c.fullJID, c.gen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// closeWith is the registry close callback. It sends </stream:stream> best
|
||||||
|
// effort and then closes the socket.
|
||||||
|
func (c *Conn) closeWith(_ error) {
|
||||||
|
_ = c.writeRaw([]byte("</stream:stream>"))
|
||||||
|
c.close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) close() {
|
||||||
|
c.mu.Lock()
|
||||||
|
if c.closed {
|
||||||
|
c.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.closed = true
|
||||||
|
c.mu.Unlock()
|
||||||
|
_ = c.netConn.Close()
|
||||||
|
if c.pingLoop != nil {
|
||||||
|
c.pingLoop.stop()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeStanza writes a stanza or feature element, logging and emitting a
|
||||||
|
// diagnostic copy. It never logs SASL material because auth is never sent out.
|
||||||
|
func (c *Conn) writeStanza(b []byte) error {
|
||||||
|
logStanza("out", b)
|
||||||
|
c.emitDiag(session.DirectionOut, "stanza", b, "")
|
||||||
|
return c.writeRaw(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeRaw performs a mutex-protected write without logging.
|
||||||
|
func (c *Conn) writeRaw(b []byte) error {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
if c.closed {
|
||||||
|
return net.ErrClosed
|
||||||
|
}
|
||||||
|
_, err := c.netConn.Write(b)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) emitDiag(dir, kind string, xml []byte, reason string) {
|
||||||
|
if c.diag == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.diag(session.Diagnostic{
|
||||||
|
Direction: dir,
|
||||||
|
Kind: kind,
|
||||||
|
Generation: c.gen,
|
||||||
|
XML: append([]byte(nil), xml...),
|
||||||
|
Reason: reason,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func rootSpaceLocal(b []byte) (string, string, error) {
|
||||||
|
d := xml.NewDecoder(bytes.NewReader(b))
|
||||||
|
tok, err := d.Token()
|
||||||
|
if err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
start, ok := tok.(xml.StartElement)
|
||||||
|
if !ok {
|
||||||
|
return "", "", fmt.Errorf("first token is not a start element")
|
||||||
|
}
|
||||||
|
return start.Name.Space, start.Name.Local, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newStreamID() string {
|
||||||
|
b := make([]byte, 16)
|
||||||
|
_, _ = rand.Read(b)
|
||||||
|
return hex.EncodeToString(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func escapeXMLAttr(s string) string {
|
||||||
|
s = strings.ReplaceAll(s, "&", "&")
|
||||||
|
s = strings.ReplaceAll(s, "<", "<")
|
||||||
|
s = strings.ReplaceAll(s, ">", ">")
|
||||||
|
s = strings.ReplaceAll(s, "\"", """)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
type saslAuth struct {
|
||||||
|
XMLName xml.Name `xml:"urn:ietf:params:xml:ns:xmpp-sasl auth"`
|
||||||
|
Mechanism string `xml:"mechanism,attr"`
|
||||||
|
Chardata string `xml:",chardata"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type iqBind struct {
|
||||||
|
XMLName xml.Name `xml:"iq"`
|
||||||
|
ID string `xml:"id,attr"`
|
||||||
|
Type string `xml:"type,attr"`
|
||||||
|
Bind struct {
|
||||||
|
XMLName xml.Name `xml:"bind"`
|
||||||
|
Resource string `xml:"resource"`
|
||||||
|
} `xml:"bind"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type iqSession struct {
|
||||||
|
XMLName xml.Name `xml:"iq"`
|
||||||
|
ID string `xml:"id,attr"`
|
||||||
|
Type string `xml:"type,attr"`
|
||||||
|
Session struct {
|
||||||
|
XMLName xml.Name `xml:"session"`
|
||||||
|
} `xml:"session"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type presenceHello struct {
|
||||||
|
XMLName xml.Name `xml:"presence"`
|
||||||
|
Status string `xml:"status"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type iqPing struct {
|
||||||
|
XMLName xml.Name `xml:"iq"`
|
||||||
|
ID string `xml:"id,attr"`
|
||||||
|
Type string `xml:"type,attr"`
|
||||||
|
From string `xml:"from,attr"`
|
||||||
|
To string `xml:"to,attr"`
|
||||||
|
Ping struct {
|
||||||
|
XMLName xml.Name `xml:"ping"`
|
||||||
|
} `xml:"ping"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type iqResult struct {
|
||||||
|
XMLName xml.Name `xml:"iq"`
|
||||||
|
ID string `xml:"id,attr"`
|
||||||
|
Type string `xml:"type,attr"`
|
||||||
|
From string `xml:"from,attr"`
|
||||||
|
To string `xml:"to,attr"`
|
||||||
|
}
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// PingPeriod is the interval between controller announce/keepalive pings.
|
||||||
|
PingPeriod = 60 * time.Second
|
||||||
|
// PingResultTimeout is how long to wait for a matching ping result before
|
||||||
|
// treating the session as half-open.
|
||||||
|
PingResultTimeout = 12 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
var pingCounter atomic.Uint64
|
||||||
|
|
||||||
|
// pingLoop sends periodic controller pings and waits for matching results.
|
||||||
|
type pingLoop struct {
|
||||||
|
c *Conn
|
||||||
|
mu sync.Mutex
|
||||||
|
pending string // current outstanding ping id
|
||||||
|
announced bool // AnnounceOK already emitted for this connection
|
||||||
|
resultCh chan string
|
||||||
|
stopCh chan struct{}
|
||||||
|
doneCh chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newPingLoop(c *Conn) *pingLoop {
|
||||||
|
return &pingLoop{
|
||||||
|
c: c,
|
||||||
|
resultCh: make(chan string, 1),
|
||||||
|
stopCh: make(chan struct{}),
|
||||||
|
doneCh: make(chan struct{}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *pingLoop) nextID() string {
|
||||||
|
return fmt.Sprintf("%d", pingCounter.Add(1))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *pingLoop) run() {
|
||||||
|
defer close(p.doneCh)
|
||||||
|
send := func(b []byte) error { return p.c.registry.Send(p.c.fullJID, p.c.gen, b) }
|
||||||
|
first := true
|
||||||
|
for {
|
||||||
|
if !first {
|
||||||
|
select {
|
||||||
|
case <-p.stopCh:
|
||||||
|
return
|
||||||
|
case <-p.c.clock.After(PingPeriod):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
first = false
|
||||||
|
|
||||||
|
id := p.nextID()
|
||||||
|
p.mu.Lock()
|
||||||
|
p.pending = id
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
|
ping := fmt.Sprintf(`<iq id="%s" to="%s" from="%s" type="get"><ping xmlns="urn:xmpp:ping"/></iq>`,
|
||||||
|
id, escapeXMLAttr(p.c.fullJID), escapeXMLAttr(p.c.cfg.ControllerJID))
|
||||||
|
if err := send([]byte(ping)); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
timeout := p.c.clock.After(PingResultTimeout)
|
||||||
|
waitResult:
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-p.stopCh:
|
||||||
|
return
|
||||||
|
case <-timeout:
|
||||||
|
p.mu.Lock()
|
||||||
|
still := p.pending == id
|
||||||
|
p.mu.Unlock()
|
||||||
|
if still {
|
||||||
|
p.c.endSession(session.ReasonPingTimeout)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
case rid := <-p.resultCh:
|
||||||
|
p.mu.Lock()
|
||||||
|
if p.pending != rid {
|
||||||
|
p.mu.Unlock()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
p.pending = ""
|
||||||
|
announced := p.announced
|
||||||
|
p.announced = true
|
||||||
|
p.mu.Unlock()
|
||||||
|
if !announced {
|
||||||
|
p.c.bus.AnnounceOK(session.ReadyEvent{
|
||||||
|
Generation: p.c.gen,
|
||||||
|
JID: p.c.fullJID,
|
||||||
|
Serial: p.c.serial,
|
||||||
|
Send: send,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
break waitResult
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// result reports a matching ping result id. It returns true when the id is the
|
||||||
|
// currently outstanding ping.
|
||||||
|
func (p *pingLoop) result(id string) bool {
|
||||||
|
p.mu.Lock()
|
||||||
|
pending := p.pending == id
|
||||||
|
p.mu.Unlock()
|
||||||
|
if !pending {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case p.resultCh <- id:
|
||||||
|
return true
|
||||||
|
case <-p.doneCh:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// stop signals the ping loop to exit.
|
||||||
|
func (p *pingLoop) stop() {
|
||||||
|
select {
|
||||||
|
case <-p.stopCh:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
close(p.stopCh)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
)
|
||||||
|
|
||||||
|
const saslNS = "urn:ietf:params:xml:ns:xmpp-sasl"
|
||||||
|
|
||||||
|
// Redact returns a copy of xml with the character data of any <auth> element
|
||||||
|
// in the SASL namespace replaced by the literal string "redacted". It is used
|
||||||
|
// before logging or diagnostic publication.
|
||||||
|
func Redact(xml []byte) []byte {
|
||||||
|
var out []byte
|
||||||
|
i := 0
|
||||||
|
for i < len(xml) {
|
||||||
|
j := bytes.Index(xml[i:], []byte("<auth"))
|
||||||
|
if j < 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
j += i
|
||||||
|
endTag, ok := scanTagEnd(xml[j:])
|
||||||
|
if !ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
startTag := xml[j : j+endTag]
|
||||||
|
isSASL := bytes.Contains(startTag, []byte(`xmlns="`+saslNS+`"`)) ||
|
||||||
|
bytes.Contains(startTag, []byte(`xmlns='`+saslNS+`'`))
|
||||||
|
if !isSASL {
|
||||||
|
out = append(out, xml[i:j+endTag]...)
|
||||||
|
i = j + endTag
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
closeStart := bytes.Index(xml[j+endTag:], []byte("</auth>"))
|
||||||
|
if closeStart < 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
closeStart += j + endTag
|
||||||
|
out = append(out, xml[i:j+endTag]...)
|
||||||
|
out = append(out, []byte("redacted")...)
|
||||||
|
out = append(out, xml[closeStart:closeStart+len("</auth>")]...)
|
||||||
|
i = closeStart + len("</auth>")
|
||||||
|
}
|
||||||
|
out = append(out, xml[i:]...)
|
||||||
|
return out
|
||||||
|
}
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/base64"
|
||||||
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
var authcidRe = regexp.MustCompile(`^[A-Za-z0-9]{1,64}$`)
|
||||||
|
|
||||||
|
// SASLAuth parses a PLAIN mechanism auth element value.
|
||||||
|
type SASLAuth struct {
|
||||||
|
Authzid string
|
||||||
|
Authcid string // robot serial from the base64 credential
|
||||||
|
Password string
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseSASLPlain decodes a base64 PLAIN token. It accepts an empty password but
|
||||||
|
// rejects missing fields or an invalid authcid.
|
||||||
|
func ParseSASLPlain(token string) (SASLAuth, error) {
|
||||||
|
decoded, err := base64.StdEncoding.DecodeString(token)
|
||||||
|
if err != nil {
|
||||||
|
return SASLAuth{}, fmt.Errorf("base64: %w", err)
|
||||||
|
}
|
||||||
|
parts := strings.Split(string(decoded), "\x00")
|
||||||
|
if len(parts) != 3 {
|
||||||
|
return SASLAuth{}, fmt.Errorf("expected 3 NUL-separated fields, got %d", len(parts))
|
||||||
|
}
|
||||||
|
auth := SASLAuth{Authzid: parts[0], Authcid: parts[1], Password: parts[2]}
|
||||||
|
if auth.Authcid == "" || !authcidRe.MatchString(auth.Authcid) {
|
||||||
|
return SASLAuth{}, fmt.Errorf("invalid authcid")
|
||||||
|
}
|
||||||
|
return auth, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SASLSuccessXML is the SASL success stanza.
|
||||||
|
var SASLSuccessXML = []byte(`<success xmlns="urn:ietf:params:xml:ns:xmpp-sasl"/>`)
|
||||||
|
|
||||||
|
// SASLFailureXML returns a failure stanza for the named condition.
|
||||||
|
func SASLFailureXML(condition string) []byte {
|
||||||
|
return []byte(fmt.Sprintf(`<failure xmlns="urn:ietf:params:xml:ns:xmpp-sasl"><%s/></failure>`, condition))
|
||||||
|
}
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Clock abstracts time for the ping loop so tests can advance the clock without
|
||||||
|
// sleeping.
|
||||||
|
type Clock interface {
|
||||||
|
Now() time.Time
|
||||||
|
After(d time.Duration) <-chan time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type realClock struct{}
|
||||||
|
|
||||||
|
func (realClock) Now() time.Time { return time.Now() }
|
||||||
|
func (realClock) After(d time.Duration) <-chan time.Time { return time.After(d) }
|
||||||
|
|
||||||
|
// Server holds the runtime dependencies for the XMPP listener.
|
||||||
|
type Server struct {
|
||||||
|
cfg config.Config
|
||||||
|
registry *session.Registry
|
||||||
|
bus *session.Bus
|
||||||
|
clock Clock
|
||||||
|
diag session.DiagnosticSink
|
||||||
|
ln net.Listener
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewServer creates a server. If clock is nil, real time is used.
|
||||||
|
func NewServer(cfg config.Config, registry *session.Registry, bus *session.Bus, clock Clock, diag session.DiagnosticSink) *Server {
|
||||||
|
if clock == nil {
|
||||||
|
clock = realClock{}
|
||||||
|
}
|
||||||
|
return &Server{
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: clock,
|
||||||
|
diag: diag,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serve starts the plaintext XMPP listener and blocks until ctx is cancelled.
|
||||||
|
// It returns nil on graceful shutdown and the listen error on failure.
|
||||||
|
func (s *Server) Serve(ctx context.Context) error {
|
||||||
|
addr := net.JoinHostPort(s.cfg.BindAddress, fmt.Sprintf("%d", s.cfg.PortXMPP))
|
||||||
|
ln, err := net.Listen("tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s: %w", addr, err)
|
||||||
|
}
|
||||||
|
s.ln = ln
|
||||||
|
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
errCh <- s.acceptLoop(ctx)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
_ = ln.Close()
|
||||||
|
return nil
|
||||||
|
case err := <-errCh:
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) acceptLoop(ctx context.Context) error {
|
||||||
|
for {
|
||||||
|
conn, err := s.ln.Accept()
|
||||||
|
if err != nil {
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c := &Conn{
|
||||||
|
netConn: conn,
|
||||||
|
cfg: s.cfg,
|
||||||
|
registry: s.registry,
|
||||||
|
bus: s.bus,
|
||||||
|
clock: s.clock,
|
||||||
|
diag: s.diag,
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Shutdown writes </stream:stream> best-effort on every current generation and
|
||||||
|
// closes the underlying sockets. It is safe to call on a nil Server.
|
||||||
|
func (s *Server) Shutdown(ctx context.Context) error {
|
||||||
|
if s == nil || s.registry == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
s.registry.CloseAll()
|
||||||
|
if s.ln != nil {
|
||||||
|
_ = s.ln.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// logStanza logs a redacted copy of a stanza at debug level.
|
||||||
|
func logStanza(dir string, b []byte) {
|
||||||
|
slog.Debug("xmpp", "dir", dir, "stanza", string(Redact(b)))
|
||||||
|
}
|
||||||
@@ -0,0 +1,341 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Tokenizer frames XMPP stanzas from a byte stream. It does not use an
|
||||||
|
// incremental XML decoder because the robot stream is intentionally incomplete:
|
||||||
|
// <stream:stream> stays open, a stanza can span TCP reads, and several stanzas
|
||||||
|
// can share one TCP segment.
|
||||||
|
type Tokenizer struct {
|
||||||
|
state tokState
|
||||||
|
buf []byte
|
||||||
|
pos int // current scan position within buf
|
||||||
|
depth int
|
||||||
|
stanzaStart int
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
type tokState int
|
||||||
|
|
||||||
|
const (
|
||||||
|
stateNeedStream tokState = iota
|
||||||
|
stateInStream
|
||||||
|
)
|
||||||
|
|
||||||
|
// EventKind identifies the kind of token emitted by the tokenizer.
|
||||||
|
type EventKind int
|
||||||
|
|
||||||
|
const (
|
||||||
|
StreamOpen EventKind = iota
|
||||||
|
Stanza
|
||||||
|
StreamClose
|
||||||
|
TokenError
|
||||||
|
)
|
||||||
|
|
||||||
|
// Event is one token from the XMPP byte stream.
|
||||||
|
type Event struct {
|
||||||
|
Kind EventKind
|
||||||
|
Data []byte
|
||||||
|
Attrs map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTokenizer returns an empty tokenizer in the need-stream state.
|
||||||
|
func NewTokenizer() *Tokenizer {
|
||||||
|
return &Tokenizer{state: stateNeedStream}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExpectNewStream resets the tokenizer to accept a fresh XML declaration and
|
||||||
|
// stream open without requiring a closing </stream:stream>. This is used after
|
||||||
|
// a SASL success.
|
||||||
|
func (t *Tokenizer) ExpectNewStream() {
|
||||||
|
t.state = stateNeedStream
|
||||||
|
t.depth = 0
|
||||||
|
t.stanzaStart = 0
|
||||||
|
t.pos = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Feed adds bytes and returns any complete events produced. After a TokenError
|
||||||
|
// event, no further events are emitted until Reset/ExpectNewStream is called.
|
||||||
|
func (t *Tokenizer) Feed(p []byte) []Event {
|
||||||
|
if t.err != nil && t.state != stateNeedStream {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
t.buf = append(t.buf, p...)
|
||||||
|
var events []Event
|
||||||
|
for {
|
||||||
|
ev, ok := t.advance()
|
||||||
|
if !ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
events = append(events, ev)
|
||||||
|
if ev.Kind == TokenError {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return events
|
||||||
|
}
|
||||||
|
|
||||||
|
// Err returns the terminal framing error, if any.
|
||||||
|
func (t *Tokenizer) Err() error { return t.err }
|
||||||
|
|
||||||
|
func (t *Tokenizer) advance() (Event, bool) {
|
||||||
|
switch t.state {
|
||||||
|
case stateNeedStream:
|
||||||
|
return t.advanceNeedStream()
|
||||||
|
case stateInStream:
|
||||||
|
return t.advanceInStream()
|
||||||
|
}
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Tokenizer) advanceNeedStream() (Event, bool) {
|
||||||
|
for {
|
||||||
|
t.pos += skipSpace(t.buf[t.pos:])
|
||||||
|
if t.pos >= len(t.buf) {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
if bytes.HasPrefix(t.buf[t.pos:], []byte("<?xml")) {
|
||||||
|
end, ok := scanPIEnd(t.buf[t.pos:])
|
||||||
|
if !ok {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
t.pos += end
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if !bytes.HasPrefix(t.buf[t.pos:], []byte("<stream:stream")) {
|
||||||
|
if len(t.buf[t.pos:]) < len("<stream:stream") {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
t.err = fmt.Errorf("expected stream open, got %q", firstToken(t.buf[t.pos:]))
|
||||||
|
return Event{Kind: TokenError, Data: copyBytes(t.buf[t.pos:])}, true
|
||||||
|
}
|
||||||
|
end, ok := scanTagEnd(t.buf[t.pos:])
|
||||||
|
if !ok {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
tag := t.buf[t.pos : t.pos+end]
|
||||||
|
attrs := parseAttrs(tag)
|
||||||
|
t.buf = t.buf[t.pos+end:]
|
||||||
|
t.pos = 0
|
||||||
|
t.state = stateInStream
|
||||||
|
t.depth = 0
|
||||||
|
t.stanzaStart = 0
|
||||||
|
return Event{Kind: StreamOpen, Data: copyBytes(tag), Attrs: attrs}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Tokenizer) advanceInStream() (Event, bool) {
|
||||||
|
for {
|
||||||
|
if t.depth == 0 {
|
||||||
|
t.pos += skipSpace(t.buf[t.pos:])
|
||||||
|
}
|
||||||
|
if t.pos >= len(t.buf) {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
if t.buf[t.pos] != '<' {
|
||||||
|
if t.depth == 0 {
|
||||||
|
t.err = fmt.Errorf("unexpected text at stream level")
|
||||||
|
t.buf = t.buf[t.pos:]
|
||||||
|
return Event{Kind: TokenError, Data: copyBytes(t.buf)}, true
|
||||||
|
}
|
||||||
|
// Inside a stanza we skip text content until the next tag.
|
||||||
|
next := bytes.IndexByte(t.buf[t.pos:], '<')
|
||||||
|
if next < 0 {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
t.pos += next
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
end, ok := scanTagEnd(t.buf[t.pos:])
|
||||||
|
if !ok {
|
||||||
|
return Event{}, false
|
||||||
|
}
|
||||||
|
tag := t.buf[t.pos : t.pos+end]
|
||||||
|
nextPos := t.pos + end
|
||||||
|
|
||||||
|
if bytes.HasPrefix(tag, []byte("</stream:stream")) {
|
||||||
|
e := Event{Kind: StreamClose, Data: copyBytes(tag)}
|
||||||
|
t.buf = t.buf[nextPos:]
|
||||||
|
t.pos = 0
|
||||||
|
t.state = stateNeedStream
|
||||||
|
t.depth = 0
|
||||||
|
t.stanzaStart = 0
|
||||||
|
return e, true
|
||||||
|
}
|
||||||
|
|
||||||
|
if bytes.HasPrefix(tag, []byte("</")) {
|
||||||
|
t.depth--
|
||||||
|
if t.depth < 0 {
|
||||||
|
t.err = fmt.Errorf("close tag without matching open")
|
||||||
|
t.buf = t.buf[nextPos:]
|
||||||
|
t.pos = 0
|
||||||
|
return Event{Kind: TokenError, Data: copyBytes(tag)}, true
|
||||||
|
}
|
||||||
|
if t.depth == 0 {
|
||||||
|
data := copyBytes(t.buf[t.stanzaStart:nextPos])
|
||||||
|
t.buf = t.buf[nextPos:]
|
||||||
|
t.pos = 0
|
||||||
|
t.stanzaStart = 0
|
||||||
|
return Event{Kind: Stanza, Data: data}, true
|
||||||
|
}
|
||||||
|
t.pos = nextPos
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
sc := isSelfClosing(tag)
|
||||||
|
if t.depth == 0 {
|
||||||
|
if sc {
|
||||||
|
data := copyBytes(t.buf[t.pos:nextPos])
|
||||||
|
t.buf = t.buf[nextPos:]
|
||||||
|
t.pos = 0
|
||||||
|
t.stanzaStart = 0
|
||||||
|
return Event{Kind: Stanza, Data: data}, true
|
||||||
|
}
|
||||||
|
t.stanzaStart = t.pos
|
||||||
|
t.depth++
|
||||||
|
t.pos = nextPos
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !sc {
|
||||||
|
t.depth++
|
||||||
|
}
|
||||||
|
t.pos = nextPos
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func skipSpace(b []byte) int {
|
||||||
|
for i, c := range b {
|
||||||
|
if c != ' ' && c != '\t' && c != '\n' && c != '\r' {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return len(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// scanTagEnd returns the index just past the matching '>' for a tag that starts
|
||||||
|
// at data[0] == '<'. It is quote-aware and returns ok=false if the tag is not
|
||||||
|
// yet complete.
|
||||||
|
func scanTagEnd(data []byte) (int, bool) {
|
||||||
|
var inQuote byte
|
||||||
|
for i := 1; i < len(data); i++ {
|
||||||
|
c := data[i]
|
||||||
|
if inQuote != 0 {
|
||||||
|
if c == inQuote {
|
||||||
|
inQuote = 0
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if c == '"' || c == '\'' {
|
||||||
|
inQuote = c
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if c == '>' {
|
||||||
|
return i + 1, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// scanPIEnd returns the index just past the matching '?>' for a processing
|
||||||
|
// instruction that starts with '<?'.
|
||||||
|
func scanPIEnd(data []byte) (int, bool) {
|
||||||
|
var inQuote byte
|
||||||
|
for i := 2; i < len(data); i++ {
|
||||||
|
c := data[i]
|
||||||
|
if inQuote != 0 {
|
||||||
|
if c == inQuote {
|
||||||
|
inQuote = 0
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if c == '"' || c == '\'' {
|
||||||
|
inQuote = c
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if c == '?' && i+1 < len(data) && data[i+1] == '>' {
|
||||||
|
return i + 2, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func isSelfClosing(tag []byte) bool {
|
||||||
|
i := len(tag) - 2 // position before '>'
|
||||||
|
for i >= 0 && isSpace(tag[i]) {
|
||||||
|
i--
|
||||||
|
}
|
||||||
|
return i >= 0 && tag[i] == '/'
|
||||||
|
}
|
||||||
|
|
||||||
|
func isSpace(c byte) bool {
|
||||||
|
return c == ' ' || c == '\t' || c == '\n' || c == '\r'
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseAttrs(tag []byte) map[string]string {
|
||||||
|
attrs := make(map[string]string)
|
||||||
|
i := 1
|
||||||
|
// skip element name
|
||||||
|
for i < len(tag) && !isSpace(tag[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
for i < len(tag) {
|
||||||
|
for i < len(tag) && isSpace(tag[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
if i >= len(tag) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
nameStart := i
|
||||||
|
for i < len(tag) && tag[i] != '=' && !isSpace(tag[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
name := string(tag[nameStart:i])
|
||||||
|
for i < len(tag) && isSpace(tag[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
if i >= len(tag) || tag[i] != '=' {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
i++ // skip '='
|
||||||
|
for i < len(tag) && isSpace(tag[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
if i >= len(tag) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
quote := tag[i]
|
||||||
|
if quote != '"' && quote != '\'' {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
valStart := i
|
||||||
|
for i < len(tag) && tag[i] != quote {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
if i >= len(tag) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
attrs[name] = string(tag[valStart:i])
|
||||||
|
i++ // skip closing quote
|
||||||
|
}
|
||||||
|
return attrs
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyBytes(b []byte) []byte {
|
||||||
|
out := make([]byte, len(b))
|
||||||
|
copy(out, b)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstToken(b []byte) []byte {
|
||||||
|
end := bytes.IndexAny(b, " \t\n\r><")
|
||||||
|
if end < 0 {
|
||||||
|
return copyBytes(b)
|
||||||
|
}
|
||||||
|
return copyBytes(b[:end])
|
||||||
|
}
|
||||||
@@ -0,0 +1,778 @@
|
|||||||
|
package xmpp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
|
||||||
|
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
|
||||||
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
// Keep tests quiet unless they explicitly fail.
|
||||||
|
slog.SetDefault(slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError})))
|
||||||
|
}
|
||||||
|
|
||||||
|
func testConfig() config.Config {
|
||||||
|
cfg, err := config.Load([]string{
|
||||||
|
"ADVERTISE_IP=192.0.2.10",
|
||||||
|
"MQTT_HOST=mqtt.example.invalid",
|
||||||
|
"CONTROLLER_JID=n95bridge@ecouser.net/homeassistant",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenizerSeveralStanzasOneSegment(t *testing.T) {
|
||||||
|
tok := NewTokenizer()
|
||||||
|
// Feed the stream open first so the tokenizer is in-stream.
|
||||||
|
data := `<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>` +
|
||||||
|
`<iq from='serial@155.ecorobot.net/atom' to='155.ecorobot.net' id='1234567890' type='get'><ping xmlns='urn:xmpp:ping'/></iq>` +
|
||||||
|
`<iq type='set' id='a'><bind xmlns='urn:ietf:params:xml:ns:xmpp-bind'><resource>atom</resource></bind></iq>` +
|
||||||
|
`<presence><status>hello world</status></presence>`
|
||||||
|
evs := tok.Feed([]byte(data))
|
||||||
|
|
||||||
|
var opens, stanzas int
|
||||||
|
for _, ev := range evs {
|
||||||
|
switch ev.Kind {
|
||||||
|
case StreamOpen:
|
||||||
|
opens++
|
||||||
|
case Stanza:
|
||||||
|
stanzas++
|
||||||
|
case TokenError:
|
||||||
|
t.Fatalf("tokenizer error: %v, data so far %q", tok.Err(), ev.Data)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if opens != 1 {
|
||||||
|
t.Fatalf("opens = %d, want 1", opens)
|
||||||
|
}
|
||||||
|
if stanzas != 3 {
|
||||||
|
t.Fatalf("stanzas = %d, want 3", stanzas)
|
||||||
|
}
|
||||||
|
if tok.Err() != nil {
|
||||||
|
t.Fatalf("tokenizer error after feed: %v", tok.Err())
|
||||||
|
}
|
||||||
|
// Remaining buffer should be empty after the three complete stanzas.
|
||||||
|
if len(tok.buf) != 0 {
|
||||||
|
t.Fatalf("remaining buffer = %q, want empty", tok.buf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenizerStanzaSpansSegments(t *testing.T) {
|
||||||
|
tok := NewTokenizer()
|
||||||
|
tok.Feed([]byte(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`))
|
||||||
|
part1 := `<iq from='serial@155.ecorobot.net/atom' to='155.ecorobot.net' id='1000000000' type='get'><ping xmlns='urn:xmpp:ping'/></iq><iq from='serial@155.ecorobot.net/atom' to='155.ecorobot.net' id='1000000001' type='get'`
|
||||||
|
part2 := `><ping xmlns='urn:xmpp:ping'/></iq>`
|
||||||
|
evs1 := tok.Feed([]byte(part1))
|
||||||
|
if len(evs1) != 1 {
|
||||||
|
t.Fatalf("first feed events = %d, want 1: %+v", len(evs1), evs1)
|
||||||
|
}
|
||||||
|
evs2 := tok.Feed([]byte(part2))
|
||||||
|
if len(evs2) != 1 {
|
||||||
|
t.Fatalf("second feed events = %d, want 1: %+v", len(evs2), evs2)
|
||||||
|
}
|
||||||
|
if evs2[0].Kind != Stanza {
|
||||||
|
t.Fatalf("second event kind = %v, want Stanza", evs2[0].Kind)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStreamIDRepeatedAfterSASL(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
rec := &recorder{}
|
||||||
|
bus.Register(rec)
|
||||||
|
|
||||||
|
cfg := testConfig()
|
||||||
|
clock := newFakeClock()
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: clock,
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
xc.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
open1 := xc.recvUntil("</stream:features>")
|
||||||
|
id1 := extractAttr(open1, "id")
|
||||||
|
if id1 == "" {
|
||||||
|
t.Fatalf("first stream id missing: %q", open1)
|
||||||
|
}
|
||||||
|
|
||||||
|
xc.send(saslPlain("serial1", "secret"))
|
||||||
|
xc.recvUntil("<success")
|
||||||
|
|
||||||
|
xc.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
open2 := xc.recvUntil("</stream:features>")
|
||||||
|
id2 := extractAttr(open2, "id")
|
||||||
|
if id1 != id2 {
|
||||||
|
t.Fatalf("stream id changed: %q -> %q", id1, id2)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A second TCP connection must receive a different id.
|
||||||
|
client2, server2 := net.Pipe()
|
||||||
|
defer client2.Close()
|
||||||
|
defer server2.Close()
|
||||||
|
c2 := &Conn{
|
||||||
|
netConn: server2,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: session.NewRegistry(),
|
||||||
|
bus: session.NewBus(),
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c2.serve()
|
||||||
|
xc2 := &xmppClient{t: t, c: client2}
|
||||||
|
xc2.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
open3 := xc2.recvUntil("</stream:features>")
|
||||||
|
id3 := extractAttr(open3, "id")
|
||||||
|
if id3 == "" || id3 == id1 {
|
||||||
|
t.Fatalf("second connection stream id not unique: %q", id3)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStartTLSAdvertisedNotWaited(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: testConfig(),
|
||||||
|
registry: session.NewRegistry(),
|
||||||
|
bus: session.NewBus(),
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
xc.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
open1 := xc.recvUntil("</stream:features>")
|
||||||
|
if !strings.Contains(open1, `<starttls xmlns="urn:ietf:params:xml:ns:xmpp-tls"><required/></starttls>`) {
|
||||||
|
t.Fatalf("missing starttls required: %q", open1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip STARTTLS and send SASL directly; server answers success.
|
||||||
|
xc.send(saslPlain("serial2", "secret"))
|
||||||
|
xc.recvUntil("<success")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandshakeToReady(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
rec := &recorder{}
|
||||||
|
bus.Register(rec)
|
||||||
|
|
||||||
|
cfg := testConfig()
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
jid := completeHandshake(t, client, "serial3", "atom")
|
||||||
|
wantJID := "serial3@155.ecorobot.net/atom"
|
||||||
|
if jid != wantJID {
|
||||||
|
t.Fatalf("bind JID = %q, want %q", jid, wantJID)
|
||||||
|
}
|
||||||
|
|
||||||
|
mustWaitFor(t, func() bool {
|
||||||
|
rec.mu.Lock()
|
||||||
|
defer rec.mu.Unlock()
|
||||||
|
return len(rec.ready) == 1
|
||||||
|
})
|
||||||
|
rec.mu.Lock()
|
||||||
|
readyCount := len(rec.ready)
|
||||||
|
rec.mu.Unlock()
|
||||||
|
if readyCount != 1 {
|
||||||
|
t.Fatalf("SessionReady count = %d, want 1", readyCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
currentGen, ok := registry.Current(wantJID)
|
||||||
|
if !ok || currentGen != 1 {
|
||||||
|
t.Fatalf("Current(%s) = %d %v, want 1 true", wantJID, currentGen, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDummyPresenceSpaces(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: testConfig(),
|
||||||
|
registry: session.NewRegistry(),
|
||||||
|
bus: session.NewBus(),
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
jid, dummy := completeHandshakeEx(t, client, "serial4", "atom")
|
||||||
|
wantTo := fmt.Sprintf(`to="%s"`, jid)
|
||||||
|
if !strings.Contains(dummy, `<presence`) || !strings.Contains(dummy, wantTo) || !strings.Contains(dummy, "> dummy </presence>") {
|
||||||
|
t.Fatalf("dummy presence = %q, want to=%q with > dummy </presence>", dummy, jid)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSASLAcceptedAndNotLogged(t *testing.T) {
|
||||||
|
// Capture all log records.
|
||||||
|
var logBuf bytes.Buffer
|
||||||
|
h := &captureHandler{Level: slog.LevelDebug}
|
||||||
|
slog.SetDefault(slog.New(h))
|
||||||
|
defer slog.SetDefault(slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError})))
|
||||||
|
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: testConfig(),
|
||||||
|
registry: session.NewRegistry(),
|
||||||
|
bus: session.NewBus(),
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
xc.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
xc.recvUntil("</stream:features>")
|
||||||
|
|
||||||
|
// Valid PLAIN credential.
|
||||||
|
payload := "\x00serial5\x00any-password"
|
||||||
|
token := base64.StdEncoding.EncodeToString([]byte(payload))
|
||||||
|
xc.send(fmt.Sprintf(`<auth xmlns='urn:ietf:params:xml:ns:xmpp-sasl' mechanism='PLAIN'>%s</auth>`, token))
|
||||||
|
xc.recvUntil("<success")
|
||||||
|
_ = logBuf.String()
|
||||||
|
|
||||||
|
h.mu.Lock()
|
||||||
|
records := h.records
|
||||||
|
h.mu.Unlock()
|
||||||
|
joined := strings.Join(records, "\n")
|
||||||
|
if strings.Contains(joined, token) {
|
||||||
|
t.Fatalf("log contained base64 token")
|
||||||
|
}
|
||||||
|
if strings.Contains(joined, "any-password") {
|
||||||
|
t.Fatalf("log contained password")
|
||||||
|
}
|
||||||
|
if !strings.Contains(joined, "serial5") {
|
||||||
|
t.Fatalf("log did not contain authcid: %q", joined)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Malformed credential.
|
||||||
|
client2, server2 := net.Pipe()
|
||||||
|
defer client2.Close()
|
||||||
|
defer server2.Close()
|
||||||
|
h2 := &captureHandler{Level: slog.LevelDebug}
|
||||||
|
slog.SetDefault(slog.New(h2))
|
||||||
|
c2 := &Conn{
|
||||||
|
netConn: server2,
|
||||||
|
cfg: testConfig(),
|
||||||
|
registry: session.NewRegistry(),
|
||||||
|
bus: session.NewBus(),
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c2.serve()
|
||||||
|
xc2 := &xmppClient{t: t, c: client2}
|
||||||
|
xc2.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
xc2.recvUntil("</stream:features>")
|
||||||
|
xc2.send(`<auth xmlns='urn:ietf:params:xml:ns:xmpp-sasl' mechanism='PLAIN'>!!!</auth>`)
|
||||||
|
fail := xc2.recvUntil("</failure>")
|
||||||
|
if !strings.Contains(fail, "malformed-request") {
|
||||||
|
t.Fatalf("expected malformed-request failure: %q", fail)
|
||||||
|
}
|
||||||
|
h2.mu.Lock()
|
||||||
|
records2 := h2.records
|
||||||
|
h2.mu.Unlock()
|
||||||
|
joined2 := strings.Join(records2, "\n")
|
||||||
|
if strings.Contains(joined2, "!!!") {
|
||||||
|
t.Fatalf("malformed log contained raw payload")
|
||||||
|
}
|
||||||
|
if !strings.Contains(joined2, "sasl malformed") {
|
||||||
|
t.Fatalf("malformed log did not contain 'sasl malformed': %q", joined2)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSessionReplaceAtBind(t *testing.T) {
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
rec := &recorder{}
|
||||||
|
bus.Register(rec)
|
||||||
|
cfg := testConfig()
|
||||||
|
clock := newFakeClock()
|
||||||
|
|
||||||
|
// Connection A completes handshake through session, not yet READY.
|
||||||
|
clientA, serverA := net.Pipe()
|
||||||
|
defer clientA.Close()
|
||||||
|
cA := &Conn{
|
||||||
|
netConn: serverA,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: clock,
|
||||||
|
}
|
||||||
|
go cA.serve()
|
||||||
|
jid := completeHandshakeNoReady(t, clientA, "serial6", "atom")
|
||||||
|
// Drain A so the replacement close-with-write does not block on a pipe.
|
||||||
|
go func() { _, _ = io.Copy(io.Discard, clientA) }()
|
||||||
|
|
||||||
|
// Connection B starts, finishes SASL, then sends bind for the same JID.
|
||||||
|
clientB, serverB := net.Pipe()
|
||||||
|
defer clientB.Close()
|
||||||
|
cB := &Conn{
|
||||||
|
netConn: serverB,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: clock,
|
||||||
|
}
|
||||||
|
go cB.serve()
|
||||||
|
xcB := &xmppClient{t: t, c: clientB}
|
||||||
|
xcB.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
xcB.recvUntil("</stream:features>")
|
||||||
|
xcB.send(saslPlain("serial6", "pw"))
|
||||||
|
xcB.recvUntil("<success")
|
||||||
|
xcB.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
xcB.recvUntil("</stream:features>")
|
||||||
|
xcB.send(`<iq type='set' id='bindB'><bind xmlns='urn:ietf:params:xml:ns:xmpp-bind'><resource>atom</resource></bind></iq>`)
|
||||||
|
bindResult := xcB.recvUntil("</iq>")
|
||||||
|
if !strings.Contains(bindResult, jid) {
|
||||||
|
t.Fatalf("B bind result missing JID: %q", bindResult)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A's later bytes should not be delivered; send a presence on A.
|
||||||
|
clientA.Write([]byte(`<presence><status>hello world</status></presence>`))
|
||||||
|
|
||||||
|
mustWaitFor(t, func() bool {
|
||||||
|
rec.mu.Lock()
|
||||||
|
defer rec.mu.Unlock()
|
||||||
|
return len(rec.down) == 1
|
||||||
|
})
|
||||||
|
|
||||||
|
rec.mu.Lock()
|
||||||
|
replacedReasons := []string{}
|
||||||
|
for _, d := range rec.down {
|
||||||
|
replacedReasons = append(replacedReasons, d.Reason)
|
||||||
|
}
|
||||||
|
readyCount := len(rec.ready)
|
||||||
|
rec.mu.Unlock()
|
||||||
|
if len(replacedReasons) != 1 || replacedReasons[0] != session.ReasonReplaced {
|
||||||
|
t.Fatalf("replacement down events = %v, want one replaced", replacedReasons)
|
||||||
|
}
|
||||||
|
if readyCount != 0 {
|
||||||
|
t.Fatalf("SessionReady emitted from replaced A: %d", readyCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send on A's generation returns ErrStale.
|
||||||
|
if err := cA.registry.Send(jid, 1, []byte("x")); !errors.Is(err, session.ErrStale) {
|
||||||
|
t.Fatalf("Send on gen 1 error = %v, want ErrStale", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBotPingMatchingID(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: testConfig(),
|
||||||
|
registry: session.NewRegistry(),
|
||||||
|
bus: session.NewBus(),
|
||||||
|
clock: newFakeClock(),
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
completeHandshake(t, client, "serial7", "atom")
|
||||||
|
_ = xcRecvPing(t, client) // consume announce ping
|
||||||
|
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
bigID := "12345678901234567890"
|
||||||
|
xc.send(fmt.Sprintf(`<iq from='serial7@155.ecorobot.net/atom' to='155.ecorobot.net' id='%s' type='get'><ping xmlns='urn:xmpp:ping'/></iq>`, bigID))
|
||||||
|
reply := xc.recvUntil("/>")
|
||||||
|
if !strings.Contains(reply, fmt.Sprintf(`id="%s"`, bigID)) {
|
||||||
|
t.Fatalf("reply missing id: %q", reply)
|
||||||
|
}
|
||||||
|
if !strings.Contains(reply, `type="result"`) {
|
||||||
|
t.Fatalf("reply not type=result: %q", reply)
|
||||||
|
}
|
||||||
|
if !strings.Contains(reply, `from="155.ecorobot.net"`) {
|
||||||
|
t.Fatalf("reply from wrong: %q", reply)
|
||||||
|
}
|
||||||
|
if !strings.Contains(reply, fmt.Sprintf(`to="serial7@155.ecorobot.net/atom"`)) {
|
||||||
|
t.Fatalf("reply to wrong: %q", reply)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAnnouncePingAfterReady(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
rec := &recorder{}
|
||||||
|
bus.Register(rec)
|
||||||
|
cfg := testConfig()
|
||||||
|
clock := newFakeClock()
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: clock,
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
completeHandshake(t, client, "serial8", "atom")
|
||||||
|
// The announce ping is sent immediately after READY.
|
||||||
|
ping := xcRecvPing(t, client)
|
||||||
|
waitForGoroutines() // let pingLoop reach its select
|
||||||
|
if ping == "" {
|
||||||
|
t.Fatalf("no announce ping received")
|
||||||
|
}
|
||||||
|
if !strings.Contains(ping, cfg.ControllerJID) {
|
||||||
|
t.Fatalf("ping from not controller JID: %q", ping)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reply to the announce ping.
|
||||||
|
id := extractAttr(ping, "id")
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
xc.send(fmt.Sprintf(`<iq type='result' from='serial8@155.ecorobot.net/atom' to='%s' id='%s'/>`, cfg.ControllerJID, id))
|
||||||
|
|
||||||
|
// Wait for AnnounceOK.
|
||||||
|
mustWaitFor(t, func() bool {
|
||||||
|
rec.mu.Lock()
|
||||||
|
defer rec.mu.Unlock()
|
||||||
|
return len(rec.announce) == 1
|
||||||
|
})
|
||||||
|
rec.mu.Lock()
|
||||||
|
gen := rec.announce[0].Generation
|
||||||
|
rec.mu.Unlock()
|
||||||
|
if gen != 1 {
|
||||||
|
t.Fatalf("AnnounceOK generation = %d, want 1", gen)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControllerPingDeadlineSignalsDown(t *testing.T) {
|
||||||
|
client, server := net.Pipe()
|
||||||
|
defer client.Close()
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
registry := session.NewRegistry()
|
||||||
|
bus := session.NewBus()
|
||||||
|
rec := &recorder{}
|
||||||
|
bus.Register(rec)
|
||||||
|
cfg := testConfig()
|
||||||
|
clock := newFakeClock()
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
netConn: server,
|
||||||
|
cfg: cfg,
|
||||||
|
registry: registry,
|
||||||
|
bus: bus,
|
||||||
|
clock: clock,
|
||||||
|
}
|
||||||
|
go c.serve()
|
||||||
|
|
||||||
|
completeHandshake(t, client, "serial9", "atom")
|
||||||
|
ping := xcRecvPing(t, client)
|
||||||
|
waitForGoroutines() // let pingLoop reach its select
|
||||||
|
if ping == "" {
|
||||||
|
t.Fatalf("no announce ping received")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Advance 12 seconds with no reply.
|
||||||
|
clock.Advance(PingResultTimeout)
|
||||||
|
|
||||||
|
mustWaitFor(t, func() bool {
|
||||||
|
rec.mu.Lock()
|
||||||
|
defer rec.mu.Unlock()
|
||||||
|
return len(rec.down) == 1
|
||||||
|
})
|
||||||
|
rec.mu.Lock()
|
||||||
|
reason := rec.down[0].Reason
|
||||||
|
rec.mu.Unlock()
|
||||||
|
if reason != session.ReasonPingTimeout {
|
||||||
|
t.Fatalf("down reason = %q, want ping-timeout", reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// completeHandshake performs the full client handshake up to READY and returns
|
||||||
|
// the bound full JID.
|
||||||
|
func completeHandshake(t *testing.T, client net.Conn, serial, resource string) string {
|
||||||
|
jid, _ := completeHandshakeEx(t, client, serial, resource)
|
||||||
|
return jid
|
||||||
|
}
|
||||||
|
|
||||||
|
// completeHandshakeEx returns the bound JID and the raw dummy presence response.
|
||||||
|
func completeHandshakeEx(t *testing.T, client net.Conn, serial, resource string) (jid, dummy string) {
|
||||||
|
jid = completeHandshakeNoReady(t, client, serial, resource)
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
xc.send(`<presence><status>hello world</status></presence>`)
|
||||||
|
dummy = xc.recvUntil("</presence>")
|
||||||
|
if !strings.Contains(dummy, "> dummy </presence>") {
|
||||||
|
t.Fatalf("dummy presence missing: %q", dummy)
|
||||||
|
}
|
||||||
|
return jid, dummy
|
||||||
|
}
|
||||||
|
|
||||||
|
func completeHandshakeNoReady(t *testing.T, client net.Conn, serial, resource string) string {
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
xc.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
xc.recvUntil("</stream:features>")
|
||||||
|
xc.send(saslPlain(serial, "password"))
|
||||||
|
xc.recvUntil("<success")
|
||||||
|
xc.send(`<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' xmlns='jabber:client' to='155.ecorobot.net' version='1.0'>`)
|
||||||
|
xc.recvUntil("</stream:features>")
|
||||||
|
xc.send(fmt.Sprintf(`<iq type='set' id='bind1'><bind xmlns='urn:ietf:params:xml:ns:xmpp-bind'><resource>%s</resource></bind></iq>`, resource))
|
||||||
|
bindResult := xc.recvUntil("</iq>")
|
||||||
|
jid := extractJID(bindResult)
|
||||||
|
if jid == "" {
|
||||||
|
t.Fatalf("no JID in bind result: %q", bindResult)
|
||||||
|
}
|
||||||
|
xc.send(`<iq type='set' id='sess1'><session xmlns='urn:ietf:params:xml:ns:xmpp-session'/></iq>`)
|
||||||
|
xc.recvUntil("/>")
|
||||||
|
return jid
|
||||||
|
}
|
||||||
|
|
||||||
|
func saslPlain(serial, password string) string {
|
||||||
|
payload := fmt.Sprintf("\x00%s\x00%s", serial, password)
|
||||||
|
token := base64.StdEncoding.EncodeToString([]byte(payload))
|
||||||
|
return fmt.Sprintf(`<auth xmlns='urn:ietf:params:xml:ns:xmpp-sasl' mechanism='PLAIN'>%s</auth>`, token)
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractJID(s string) string {
|
||||||
|
start := strings.Index(s, "<jid>")
|
||||||
|
if start < 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
start += len("<jid>")
|
||||||
|
end := strings.Index(s[start:], "</jid>")
|
||||||
|
if end < 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return s[start : start+end]
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractAttr(s, name string) string {
|
||||||
|
prefix := name + "=\""
|
||||||
|
start := strings.Index(s, prefix)
|
||||||
|
if start < 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
start += len(prefix)
|
||||||
|
end := strings.Index(s[start:], "\"")
|
||||||
|
if end < 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return s[start : start+end]
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForGoroutines() { time.Sleep(20 * time.Millisecond) }
|
||||||
|
|
||||||
|
func xcRecvPing(t *testing.T, client net.Conn) string {
|
||||||
|
xc := &xmppClient{t: t, c: client}
|
||||||
|
return xc.recvUntil("</iq>")
|
||||||
|
}
|
||||||
|
|
||||||
|
func mustReadAll(c net.Conn) string {
|
||||||
|
c.SetReadDeadline(time.Now().Add(200 * time.Millisecond))
|
||||||
|
defer c.SetReadDeadline(time.Time{})
|
||||||
|
var buf bytes.Buffer
|
||||||
|
b := make([]byte, 4096)
|
||||||
|
for {
|
||||||
|
n, err := c.Read(b)
|
||||||
|
if n > 0 {
|
||||||
|
buf.Write(b[:n])
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func mustWaitFor(t *testing.T, f func() bool) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(2 * time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if f() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("condition not satisfied")
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmppClient struct {
|
||||||
|
t *testing.T
|
||||||
|
c net.Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (xc *xmppClient) send(s string) {
|
||||||
|
xc.t.Helper()
|
||||||
|
_, err := xc.c.Write([]byte(s))
|
||||||
|
if err != nil {
|
||||||
|
xc.t.Fatalf("client write: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (xc *xmppClient) recvUntil(marker string) string {
|
||||||
|
xc.t.Helper()
|
||||||
|
var buf bytes.Buffer
|
||||||
|
b := make([]byte, 1024)
|
||||||
|
xc.c.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||||
|
defer xc.c.SetReadDeadline(time.Time{})
|
||||||
|
for {
|
||||||
|
n, err := xc.c.Read(b)
|
||||||
|
if n > 0 {
|
||||||
|
buf.Write(b[:n])
|
||||||
|
}
|
||||||
|
if strings.Contains(buf.String(), marker) {
|
||||||
|
return buf.String()
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, os.ErrDeadlineExceeded) {
|
||||||
|
xc.t.Fatalf("timeout waiting for %q, got %q", marker, buf.String())
|
||||||
|
}
|
||||||
|
xc.t.Fatalf("client read: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type recorder struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
ready []session.ReadyEvent
|
||||||
|
announce []session.ReadyEvent
|
||||||
|
down []session.DownEvent
|
||||||
|
stanzas []session.StanzaEvent
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recorder) SessionReady(e session.ReadyEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.ready = append(r.ready, e)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recorder) AnnounceOK(e session.ReadyEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.announce = append(r.announce, e)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recorder) SessionDown(e session.DownEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.down = append(r.down, e)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recorder) Stanza(e session.StanzaEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.stanzas = append(r.stanzas, e)
|
||||||
|
}
|
||||||
|
|
||||||
|
type captureHandler struct {
|
||||||
|
Level slog.Leveler
|
||||||
|
mu sync.Mutex
|
||||||
|
records []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *captureHandler) Enabled(_ context.Context, level slog.Level) bool {
|
||||||
|
return level >= h.Level.Level()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *captureHandler) Handle(_ context.Context, r slog.Record) error {
|
||||||
|
var b strings.Builder
|
||||||
|
b.WriteString(r.Message)
|
||||||
|
r.Attrs(func(a slog.Attr) bool {
|
||||||
|
b.WriteString(fmt.Sprintf(" %s=%v", a.Key, a.Value))
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
h.mu.Lock()
|
||||||
|
h.records = append(h.records, b.String())
|
||||||
|
h.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *captureHandler) WithAttrs(attrs []slog.Attr) slog.Handler { return h }
|
||||||
|
func (h *captureHandler) WithGroup(name string) slog.Handler { return h }
|
||||||
|
|
||||||
|
// fakeClock is a deterministic clock for tests. It starts at a fixed base time.
|
||||||
|
type fakeClock struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
now time.Time
|
||||||
|
timers []*fakeTimer
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeTimer struct {
|
||||||
|
fire time.Time
|
||||||
|
ch chan time.Time
|
||||||
|
fired bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeClock() *fakeClock {
|
||||||
|
return &fakeClock{now: time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeClock) Now() time.Time {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
return f.now
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeClock) After(d time.Duration) <-chan time.Time {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
ch := make(chan time.Time, 1)
|
||||||
|
f.timers = append(f.timers, &fakeTimer{fire: f.now.Add(d), ch: ch})
|
||||||
|
f.fireDue()
|
||||||
|
return ch
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeClock) Advance(d time.Duration) {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
f.now = f.now.Add(d)
|
||||||
|
f.fireDue()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeClock) fireDue() {
|
||||||
|
for _, t := range f.timers {
|
||||||
|
if !t.fired && !f.now.Before(t.fire) {
|
||||||
|
t.fired = true
|
||||||
|
select {
|
||||||
|
case t.ch <- f.now:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user