diff --git a/CHANGELOG.md b/CHANGELOG.md index cdab958b..db613e1f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -46,6 +46,42 @@ accepts both, but the store flags the old spelling as deprecated page-by-page migration order, and how forms switch to the model and to JSON submit behind a flag. +### Control socket (stage 1: on-demand) + +- **The display now serves a control socket**, + `/run/ledmatrix/control.sock`. It carries versioned JSON commands, one per + line, and every command gets an answer + ([docs/IPC_CONTROL_SOCKET.md](docs/IPC_CONTROL_SOCKET.md)). + - On-demand start, stop and status are the first commands, plus `hello` + (version negotiation) and `ping`. + - Start and stop are acknowledged once the render thread has them queued. + The render thread applies them through the same handler as the file + mailbox, at its next on-demand check. On a scrolling screen that is the + next frame (the mailbox waits up to 0.25 s). On a static screen it is up + to 1 s, the same as the mailbox. + - The server's threads never touch rendering. Garbage, oversize messages + and slow or vanishing clients are answered or dropped without blocking the + display. + - New core modules: `src/ipc/contract.py`, `server.py` and `client.py`. + They are internal, not a plugin API. +- **`POST /api/v3/display/on-demand/start` and `/stop` try the socket + first.** On any failure (the display is stopped or predates the socket, a + timeout, a refusal), they write the `display_on_demand_request` mailbox + exactly as before. The response's new `transport` field says which path + was used (`"socket"` or `"mailbox"`), and `socket_error` gives the reason + for a fallback. Both paths carry the same `request_id`, so a request that + arrives both ways runs once. The mailbox, and the plugins that write it + directly, keep working for at least one more release. +- **Permissions.** The socket is `0660` and owned by the group the two + services already share (the cache directory's group, `ledmatrix` on an + installed device). On Linux the server also checks each connection's + `SO_PEERCRED`: root, the display's own user, or a member of that group. + `/run/ledmatrix` comes from the existing `RuntimeDirectory=` (#687), or the + display creates it as root under an older unit, so no installer or unit + change is needed. `LEDMATRIX_CONTROL_SOCKET` overrides the path for both + processes, or turns the socket off with `off`. A non-root dev run uses a + private per-user path under the temp directory. + ### Update channels - Devices no longer pick up every merge to `main`. A new setting, diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 304ba49c..7a0d11c3 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -41,7 +41,8 @@ each other. They share three things: | State | Where | Written by | Read by | |---|---|---|---| -| On-demand request | cache `display_on_demand_request` | web: `start_on_demand_display()` / `stop_on_demand_display()` in [`api_v3/display.py`](../web_interface/blueprints/api_v3/display.py) | display: `_poll_on_demand_requests()` | +| On-demand command | control socket `/run/ledmatrix/control.sock` ([IPC_CONTROL_SOCKET.md](IPC_CONTROL_SOCKET.md)) | web: `start_on_demand_display()` / `stop_on_demand_display()` in [`api_v3/display.py`](../web_interface/blueprints/api_v3/display.py), via [`src/ipc/client.py`](../src/ipc/client.py) | display: [`src/ipc/server.py`](../src/ipc/server.py) acks; the render thread applies it in `_poll_on_demand_requests()` | +| On-demand request (fallback) | cache `display_on_demand_request` | web, when the socket fails; four plugins write it directly | display: `_poll_on_demand_requests()` | | On-demand state | cache `display_on_demand_state` | display: `_publish_on_demand_state()` | web: `/api/v3/display/on-demand/status` | | Current screen | cache `display_current_state` | display | web: `/api/v3/display/current-status` | | Plugin errors | cache `plugin_error_snapshot` | display: `ErrorSnapshotPublisher` ([`src/error_aggregator.py`](../src/error_aggregator.py)) | web: `read_error_report()` for `/api/v3/errors/*` | @@ -55,9 +56,14 @@ each other. They share three things: | Render-loop heartbeat | `/run/ledmatrix/display-heartbeat.json` (tmpfs) | display: the render thread, via [`display_watchdog`](../src/display_watchdog.py) | web: `/api/v3/health` (`checks.display_loop`); the update health check | The on-demand start route starts `ledmatrix.service` when it is not running -(`start_service`, on by default) but never restarts a running one: the display -reads the mailbox every `ON_DEMAND_POLL_INTERVAL` (0.25s), from its dwell -sleep, its render loops and Vegas's interrupt check as well as the main loop. +(`start_service`, on by default) but never restarts a running one. The routes +send the command over the display's control socket and get an ack; when that +fails (a stopped display, one older than the socket) they write the mailbox +instead, which the display reads every `ON_DEMAND_POLL_INTERVAL` (0.25s), from +its dwell sleep, its render loops and Vegas's interrupt check as well as the +main loop. Both ways end in the same handler, `_handle_on_demand_request()`. +The socket's handlers only queue; see [IPC_CONTROL_SOCKET.md](IPC_CONTROL_SOCKET.md) +for the protocol, the permission model and the plan to retire the mailboxes. ### Web and display processes: who runs plugins diff --git a/docs/IPC_CONTROL_SOCKET.md b/docs/IPC_CONTROL_SOCKET.md new file mode 100644 index 00000000..0b84c6e1 --- /dev/null +++ b/docs/IPC_CONTROL_SOCKET.md @@ -0,0 +1,250 @@ +# Control socket (web → display) + +The display process serves a Unix socket that the web interface uses to send +it commands and get an answer back. It replaces the cache-file "mailboxes" on +the SD card one command at a time. Stage 1, described here, carries on-demand +start/stop/status. The file mailbox stays as a fallback for one release. + +| | | +|---|---| +| Socket | `/run/ledmatrix/control.sock` (tmpfs) | +| Served by | the display process ([`src/ipc/server.py`](../src/ipc/server.py)), started by `DisplayController.run()` | +| Used by | the web interface ([`src/ipc/client.py`](../src/ipc/client.py)): `POST /api/v3/display/on-demand/start` and `/stop` | +| Contract | [`src/ipc/contract.py`](../src/ipc/contract.py): messages, versions, framing and the socket path; both sides import it | +| Override | `LEDMATRIX_CONTROL_SOCKET=/some/path.sock` for both processes, or `=off` to disable it | + +## Why + +Before the socket, the web interface sent commands by writing a cache key +(`display_on_demand_request`) that the display read every 0.25 s. + +- **No acknowledgement.** The route answered "success" once the file was + written, whether or not a display was running to read it. +- **Lost requests.** The display had to read the request and then delete it. + A request written between those two steps could be thrown away (see + `_consume_on_demand_request`). The cache has no atomic claim to prevent it. +- **Fragile.** Each channel repeated its own permission, atomic-write, + staleness and in-memory-cache rules. Two of them caused bugs: a `memory_ttl` + bug ignored every on-demand request after the first for an hour, and a + stopped display was still reported as "active" for two minutes. + +The socket answers every command, carries one request per message (so nothing +can overwrite it), and belongs to the display process. If the display is not +running, the socket does not exist, and the web interface knows right away. + +## Protocol (version 1) + +**Framing.** One JSON object per line (newline-delimited JSON), UTF-8, at +most 64 KiB per line (`MAX_MESSAGE_BYTES`). Senders encode with +`ensure_ascii`, so a newline never appears inside a message. A connection +can carry several requests. Each request gets exactly one response, in order. + +**Request** + +```json +{"v": 1, "id": "5f0c…", "cmd": "on_demand.start", + "args": {"plugin_id": "clock", "mode": null, "duration": 30, "pinned": false}} +``` + +- `v` is the protocol version. +- `id` is a printable string of 1-128 characters. It is echoed back in the + response, and for on-demand commands it is also the on-demand `request_id`. +- `cmd` is a command name. +- `args` is an object. It may be omitted when a command takes no arguments. + +**Response** + +```json +{"v": 1, "id": "5f0c…", "ok": true, "result": {"accepted": true, "request_id": "5f0c…", "queued": 1}} +{"v": 1, "id": "5f0c…", "ok": false, "error": {"code": "busy", "message": "…"}} +``` + +`id` is `null` only when the request could not be parsed far enough to have +one. Clients branch on `error.code`, never on the message text. + +**Commands** + +| `cmd` | `args` | `result` | Kind | +|---|---|---|---| +| `hello` | `{versions: [int], client?: str}` | `{version, versions, commands, max_message_bytes, server}` | answered directly | +| `ping` | — | `{pong: true}` | answered directly | +| `on_demand.start` | `{plugin_id?, mode?, duration?, pinned?}` (at least one of `plugin_id` and `mode`) | ack | queued | +| `on_demand.stop` | — | ack | queued | +| `on_demand.status` | — | `{on_demand: {...}, current_mode, display_active}` | answered directly | + +`duration` is a number of seconds, or a numeric string. `0`, `null` or `""` +mean "until stopped". `pinned` must be a real boolean: the REST route has +already converted strings like `"false"` before it sends the command. The +`on_demand` object in `on_demand.status` is the same dict the display +publishes to `display_on_demand_state`. + +**Acknowledgements.** A queued command is *accepted*, not *done*. +`{"accepted": true, "request_id": …}` means the command is waiting in the +render thread's queue, and the render thread will apply it at its next +on-demand check. That is within one frame on a scrolling screen, 0.25 s +during a dwell, and up to 1 s on a static screen, whose frame loop sleeps a +second between frames. Except on a scrolling screen, where the mailbox waits +up to 0.25 s, these are the mailbox's delays too: stage 1 adds +acknowledgements, not speed. Any outcome is published as before +(`display_on_demand_state`, and `status`/`error` for a bad plugin or mode), +and it can be read with `on_demand.status`. + +**Versions.** Every request carries `v`. For any command except `hello`, a +`v` the display does not speak gets `unsupported_version`. `hello` is checked +by its `versions` list instead, and its result names the highest version both +sides share, so a client can find out what a display supports before it +relies on anything newer. Stage 1's client sends `v: 1` and falls back to the +mailbox when the display refuses it. It does not send `hello` first, which +saves a round trip. + +**Error codes:** `bad_json`, `bad_request`, `message_too_large`, +`unsupported_version`, `unknown_command`, `invalid_args`, `busy` (queue full, +or too many connections), `forbidden` (peer credentials refused), `internal`. + +Try it on a device: + +```bash +python3 - <<'EOF' +from src.ipc import client # run from the project directory +print(client.on_demand_status()) +EOF +``` + +## How the display applies a command + +The server's threads never touch rendering. A connection thread parses the +request, validates it against the contract, and then does one of two things: + +- For a command that changes the panel, it puts a `QueuedCommand` on a + bounded queue (16 entries) and answers with the ack. +- For a query, it answers from a status snapshot the display provides + (`DisplayController._control_status`). The snapshot only reads attributes. + +The render thread drains the queue in `_poll_on_demand_requests()`, the same +place it reads the mailbox, and hands each command to +`_handle_on_demand_request()`, which is the mailbox's own handler. The two +paths share all of their code: activation, the processed-id guard, error +publishing, and resuming the rotation afterwards. The 0.25 s floor on the +mailbox read does not apply to the queue, because draining it costs no disk +read. A queued command also lets `_service_pending_changes()` skip its own +floor, so a long scrolling screen or a Vegas iteration takes the command at +its next frame. + +**Exactly once.** A command and a mailbox write for the same request share +one `request_id`. If the client times out after the display queued the +command and then also writes the mailbox, the display processes the request +once. The existing `on_demand_request_id` and processed-id checks drop the +second copy. + +## Robustness + +All of this runs inside the display process, so nothing a client does may +block the render loop or crash it: + +- **Bounded connections.** Each connection gets its own daemon thread, with + at most 8 at once. One more is answered `busy` and closed. +- **Timeouts.** Each read and write times out after 2 s. A message must + arrive whole within 5 s of its first byte. An idle connection is closed + after 10 s. A slow or stuck client costs one thread for a few seconds. +- **Malformed input.** A line that is not JSON gets `bad_json`, and the + connection carries on. A line longer than 64 KiB gets `message_too_large`, + and the connection is closed, because the next message boundary cannot be + found. A client that disconnects mid-message is dropped silently. No + exception from a handler leaves the connection thread. +- **Full queue.** When the queue is full, the client gets `busy` and falls + back to the mailbox. A full queue means the render thread is stuck, and the + systemd watchdog deals with that. +- **Startup.** The server binds under a temporary name, sets the mode and the + group, then renames the socket into place, so it never appears with the + umask's permissions. It removes a stale socket (a file that nothing is + listening on). It never removes a live socket or a file that is not a + socket. `close()` removes the socket only if it is still the one this + process created. +- **Never fatal.** If the server cannot start (Windows, no `AF_UNIX`, a bind + failure, `LEDMATRIX_CONTROL_SOCKET=off`), it logs that and the display runs + as before. The web interface then uses the mailbox. + +## Security model + +The display runs as root and the web interface as the installing user (see +[PERMISSIONS.md](PERMISSIONS.md)). The socket admits exactly those two, plus +anything else in the group they share: + +1. **The directory.** `/run/ledmatrix` is created by `RuntimeDirectory=ledmatrix` + in `ledmatrix.service` (#687): root-owned, `0755`, on tmpfs, and removed + when the display stops. Under an older unit, the display creates the + directory itself as root, as it does for the heartbeat. No installer + change is needed. +2. **The socket file.** The file is `root:` with mode `0660`, + and the kernel refuses `connect()` to anyone without write permission on + it. The shared group is the cache directory's group whenever that + directory is group-writable. That is `ledmatrix` on an installed device + (`/var/cache/ledmatrix` is `root:ledmatrix 2775`), and it is the same rule + DiskCache uses for every file the two services share. Otherwise the group + is the project directory's (`get_shared_group_gid()`, which config files + use). With neither, the mode is `0600` and only root can connect. +3. **Peer credentials.** Where the kernel reports them (`SO_PEERCRED`, on + Linux), the server checks every connection again. It accepts root, the + display's own user, or a member of the shared group: the peer's primary + gid, or a supplementary group read from `/proc//status`. If `/proc` + is unreadable, it uses the group database. Any other peer gets `forbidden` + and is disconnected. This covers a socket mode that someone loosened by + hand. + +The commands are deliberately narrow. Stage 1 can start or stop on-demand +display and read its state, which anyone who can reach the web UI can already +do. Nothing on the socket runs a shell, writes a file, or names a path. + +**Development.** A display that is not root and cannot write to +`/run/ledmatrix`, such as `python3 run.py -e` from a checkout, serves the +socket at `$TMPDIR/ledmatrix-/control.sock`. That directory is private +(`0700`), and the server refuses it if another user owns it. The web +interface, run by the same user, looks there after `/run/ledmatrix`. The test +suite sets `LEDMATRIX_CONTROL_SOCKET=off` (`test/conftest.py`), so a run on a +device never touches the live display. + +## Stage plan + +1. **On-demand, with acks (this stage).** Contract, server, client. + `on_demand.start`/`stop`/`status`, `hello`, `ping`. The REST routes try the + socket first and report `transport: "socket" | "mailbox"` (plus + `socket_error` on fallback). The mailbox is unchanged, and the plugins that + write it directly (birdnet-go, mqtt-notifications, on-air, pomodoro-timer) + keep working. +2. **Commands that are restarts or polls today.** + - `brightness.set`, transient and with no `config.json` write. + - `plugin.reload`, which replaces the `restart_required` answer from #688 + with a live reload of the updated plugin on the render thread. + - `config.reload`, which applies a saved config without waiting for the 2 s + mtime poll and acks which sections changed. + - The dwell sleep and the static screen's 1 s frame sleep wait on the + queue instead of sleeping, so a command lands within milliseconds on + every kind of screen. Under WSL, with a static plugin on screen, a stop + takes 1.0 s by either path today. +3. **A state stream.** A `subscribe` command that keeps the connection open + and pushes events: mode changes, on-demand state, plugin runtime state and + the heartbeat. It replaces the polled `display_current_state`, + `plugin_runtime_snapshot` (#690) and `display-heartbeat.json` (#687) for + readers that hold a connection. The web interface relays it to its + existing SSE stream. The files remain for one release for older readers. +4. **Retire the mailboxes.** After a release in which every device has had the + socket, the web interface stops writing `display_on_demand_request`, and + the display stops polling it, logging the plugins that still write it so + they can move to an in-process `request_display()`. The other cache keys + used as messages (`plugin_error_clear_request` and the remaining + `display_*` keys) move to the socket or to tmpfs. + +## Checking it on a device + +```bash +ls -l /run/ledmatrix/control.sock # srw-rw---- root ledmatrix +sudo journalctl -u ledmatrix | grep "Control socket" +curl -s -X POST localhost:5000/api/v3/display/on-demand/start \ + -H 'Content-Type: application/json' -d '{"plugin_id":"clock","duration":20}' +# ... "transport": "socket" +``` + +If the response says `"transport": "mailbox"`, `socket_error` gives the +reason. `no_socket` means the display is stopped or predates the socket. +`refused` usually means the web user is not in the socket's group, which +takes effect when the web service restarts after the user is added. diff --git a/docs/PERMISSIONS.md b/docs/PERMISSIONS.md index 2d3d7125..ff1fd9ae 100644 --- a/docs/PERMISSIONS.md +++ b/docs/PERMISSIONS.md @@ -29,6 +29,8 @@ in again (services pick them up on restart). | `assets/` | web user | dirs `755`, files `644` | Root writes downloaded logos regardless | | `/var/cache/ledmatrix/` | `root:ledmatrix` | `2775` (setgid) | Shared cache: see below | | Cache files | creator : `ledmatrix` | `660` | | +| `/run/ledmatrix/` | `root` | `755` | tmpfs; `RuntimeDirectory=` in `ledmatrix.service`, removed when the display stops | +| `/run/ledmatrix/control.sock` | `root` : cache directory's group (`ledmatrix`) | `660` | The display's control socket; only root and that group can connect. See [IPC_CONTROL_SOCKET.md](IPC_CONTROL_SOCKET.md#security-model) | | `scripts/fix_perms/safe_plugin_rm.sh`, `safe_pip_install.sh` | `root:root` | `755` | Run as root through sudo, so the web user must not be able to edit them | | `/etc/sudoers.d/ledmatrix_web`, `ledmatrix_wifi` | `root` | `440` | | diff --git a/docs/README.md b/docs/README.md index f4536ac2..f2c664df 100644 --- a/docs/README.md +++ b/docs/README.md @@ -73,6 +73,7 @@ Going deeper: - [ARCHITECTURE.md](ARCHITECTURE.md) — processes, display loop, plugin system, web UI; where to start reading - [WEB_FRONTEND_ARCHITECTURE.md](WEB_FRONTEND_ARCHITECTURE.md) — the web UI's ES modules, page lifecycle and form model, and the page-by-page migration to them +- [IPC_CONTROL_SOCKET.md](IPC_CONTROL_SOCKET.md) — the display's control socket: protocol, security model, stage plan - [DEVELOPMENT.md](DEVELOPMENT.md) — environment setup - [HOW_TO_RUN_TESTS.md](HOW_TO_RUN_TESTS.md) — running the test suite - [MULTI_ROOT_WORKSPACE_SETUP.md](MULTI_ROOT_WORKSPACE_SETUP.md) — multi-repo workspace diff --git a/docs/REST_API_REFERENCE.md b/docs/REST_API_REFERENCE.md index b9206825..db1b0970 100644 --- a/docs/REST_API_REFERENCE.md +++ b/docs/REST_API_REFERENCE.md @@ -447,13 +447,23 @@ Request a specific plugin to display on-demand. "mode": "nfl_live", "duration": 45, "pinned": true, - "service": { "active": true, "returncode": 0, "stdout": "", "stderr": "" } + "service": { "active": true, "returncode": 0, "stdout": "", "stderr": "" }, + "transport": "socket" } } ``` `service` is `null` when `start_service` is false. +`transport` says how the request reached the display: `"socket"` means the +display's control socket acknowledged it (it is queued for the render thread; +see [IPC_CONTROL_SOCKET.md](IPC_CONTROL_SOCKET.md)), `"mailbox"` means it was +written to the cache mailbox the display polls, as before the socket existed. +With `"mailbox"`, `socket_error` gives the reason the socket was not used +(`no_socket` when the display is stopped or predates the socket, `timeout`, +`refused`, `busy`, ...). Either way the request is applied the same way; +`request_id` is the same id in both. + ### Stop On-Demand Display **POST** `/api/v3/display/on-demand/stop` @@ -476,11 +486,14 @@ Stop the current on-demand display. "status": "success", "data": { "request_id": "uuid-here", - "service": null + "service": null, + "transport": "socket" } } ``` +`transport` and `socket_error` are as for start. + --- ## Plugins diff --git a/mypy-clean.txt b/mypy-clean.txt index be7d4937..80c92425 100644 --- a/mypy-clean.txt +++ b/mypy-clean.txt @@ -50,6 +50,10 @@ src/display_geometry.py src/dynamic_team_resolver.py src/exceptions.py src/font_usage.py +src/ipc/__init__.py +src/ipc/client.py +src/ipc/contract.py +src/ipc/server.py src/logging_config.py src/logo_downloader.py src/matrix_support.py diff --git a/src/display_controller.py b/src/display_controller.py index e8704d29..ff7d5e2d 100644 --- a/src/display_controller.py +++ b/src/display_controller.py @@ -42,6 +42,7 @@ from src.cache_manager import CacheManager from src.font_manager import FontManager from src.logging_config import get_logger from src.common.sync_manager import DisplaySyncManager, SyncRole +from src.ipc.server import ControlServer, start_control_server from src.vegas_mode.render_pipeline import SYNC_SEND_INTERVAL # Get logger with consistent configuration @@ -251,6 +252,10 @@ class DisplayController: # Monotonic stamp of the last _service_pending_changes pass; same # "None means never" convention as _last_on_demand_poll. self._last_pending_service: Optional[float] = None + # The control socket (src/ipc), started by run(). None when it is not + # served (Windows, LEDMATRIX_CONTROL_SOCKET=off, a bind failure); + # the file mailbox works either way. + self._control_server = None # A brightness set_brightness() refused, so the periodic service pass # doesn't retry (and log) the same failure several times a second. self._failed_brightness_target: Optional[int] = None @@ -1313,10 +1318,12 @@ class DisplayController: def _get_on_demand_remaining(self) -> Optional[float]: """Calculate remaining time for an active on-demand session.""" - if not self.on_demand_active or self.on_demand_expires_at is None: + # Read once: the control socket's status command calls this from + # its own thread, while the render thread may be clearing the field. + expires_at = self.on_demand_expires_at + if not self.on_demand_active or expires_at is None: return None - remaining = self.on_demand_expires_at - time.time() - return max(0.0, remaining) + return max(0.0, expires_at - time.time()) def _publish_current_mode_state(self) -> None: """Publish the currently active display mode/plugin to cache for the web UI.""" @@ -1350,23 +1357,27 @@ class DisplayController: or time.monotonic() - self._last_published_at >= CURRENT_STATE_REFRESH_SECONDS): self._publish_current_mode_state() + def _on_demand_state(self) -> Dict[str, Any]: + """The on-demand state as published to the cache and the control socket.""" + return { + 'active': self.on_demand_active, + 'mode': self.on_demand_mode, + 'plugin_id': self.on_demand_plugin_id, + 'requested_at': self.on_demand_requested_at, + 'expires_at': self.on_demand_expires_at, + 'duration': self.on_demand_duration, + 'pinned': self.on_demand_pinned, + 'status': self.on_demand_status, + 'error': self.on_demand_last_error, + 'last_event': self.on_demand_last_event, + 'remaining': self._get_on_demand_remaining(), + 'last_updated': time.time() + } + def _publish_on_demand_state(self) -> None: """Publish current on-demand state to cache for external consumers.""" try: - state = { - 'active': self.on_demand_active, - 'mode': self.on_demand_mode, - 'plugin_id': self.on_demand_plugin_id, - 'requested_at': self.on_demand_requested_at, - 'expires_at': self.on_demand_expires_at, - 'duration': self.on_demand_duration, - 'pinned': self.on_demand_pinned, - 'status': self.on_demand_status, - 'error': self.on_demand_last_error, - 'last_event': self.on_demand_last_event, - 'remaining': self._get_on_demand_remaining(), - 'last_updated': time.time() - } + state = self._on_demand_state() self.cache_manager.set('display_on_demand_state', state) except (OSError, RuntimeError, ValueError, TypeError) as err: logger.error("Failed to publish on-demand state: %s", err, exc_info=True) @@ -1426,6 +1437,9 @@ class DisplayController: #: whole cost is one monotonic-clock compare. PENDING_CHANGES_INTERVAL = ON_DEMAND_POLL_INTERVAL + #: Class-level default for controllers built without __init__ (tests). + _control_server: Optional[ControlServer] = None + def _service_pending_changes(self) -> None: """Apply changes made elsewhere while the display thread is busy. @@ -1446,7 +1460,10 @@ class DisplayController: """ now = time.monotonic() last = self._last_pending_service - if last is not None and now - last < self.PENDING_CHANGES_INTERVAL: + # A command queued on the control socket skips the floor: it is in + # memory, so applying it now costs no disk read. + if (last is not None and now - last < self.PENDING_CHANGES_INTERVAL + and not (self._control_server and self._control_server.has_pending)): return self._last_pending_service = now @@ -1588,8 +1605,49 @@ class DisplayController: except (OSError, AttributeError, KeyError) as err: logger.debug("Could not clear the on-demand request mailbox: %s", err) + def _start_control_server(self) -> None: + """Serve the control socket (src/ipc/server.py). Never raises. + + Its handlers only queue commands; _drain_control_commands applies + them on the render thread, where the mailbox is read. + """ + if self._control_server is not None: + return + try: + self._control_server = start_control_server( + status_provider=self._control_status, + cache_dir=getattr(self.cache_manager, 'cache_dir', None)) + except Exception: # pylint: disable=broad-except + logger.exception("Control socket not started; using the file mailbox only") + + def _control_status(self) -> Dict[str, Any]: + """The socket's on_demand.status answer. Runs on the socket's thread: reads only.""" + return {'on_demand': self._on_demand_state(), + 'current_mode': self.current_display_mode, + 'display_active': self.is_display_active} + + def _drain_control_commands(self) -> None: + """Apply on-demand commands that arrived over the control socket. + + Each goes through _handle_on_demand_request, the mailbox's own + handler, so both ways in behave the same, and a request that came + both ways (a client that timed out and fell back) has one request + id and is processed once. + """ + server = self._control_server + if server is None or not server.has_pending: + return + for command in server.drain(): + try: + self._handle_on_demand_request(command.as_on_demand_request()) + except Exception: # pylint: disable=broad-except + logger.exception("Failed to apply control socket command %s", + command.request_id) + def _poll_on_demand_requests(self) -> None: """Poll cache for new on-demand requests from external controllers.""" + # Socket commands are already in memory: no disk read, so no floor. + self._drain_control_commands() now = time.monotonic() if (self._last_on_demand_poll is not None and now - self._last_on_demand_poll < self.ON_DEMAND_POLL_INTERVAL): @@ -1614,7 +1672,10 @@ class DisplayController: if not request: return + self._handle_on_demand_request(request) + def _handle_on_demand_request(self, request: Dict[str, Any]) -> None: + """Process one on-demand request, from the mailbox or the control socket.""" request_id = request.get('request_id') if not request_id: return @@ -2248,6 +2309,7 @@ class DisplayController: # vouch for: beats from any other thread are ignored, so a render # thread stuck inside a plugin stops them. display_watchdog.watchdog.bind_render_thread() + self._start_control_server() try: # Initialize with cached data for fast startup - let background updates refresh naturally @@ -3649,6 +3711,13 @@ class DisplayController: # First: a clean stop is not a hang, and a heartbeat left behind # would read as a frozen panel to the web interface. display_watchdog.watchdog.stopping() + # Stop taking commands; the socket file goes with it. + if self._control_server is not None: + try: + self._control_server.close() + except Exception as e: + logger.warning("Error closing the control socket: %s", e) + self._control_server = None # Stop the async update worker first so no in-flight update() call # is still touching display/cache-backed resources while they're # torn down below. diff --git a/src/ipc/__init__.py b/src/ipc/__init__.py new file mode 100644 index 00000000..571c1aab --- /dev/null +++ b/src/ipc/__init__.py @@ -0,0 +1,11 @@ +"""The display process's control socket: web -> display commands with acks. + +- :mod:`src.ipc.contract` -- the versioned messages, the framing and where + the socket lives. Shared by both sides; standard library only. +- :mod:`src.ipc.server` -- the display side: a threaded Unix-socket server + whose handlers only queue work for the render thread. +- :mod:`src.ipc.client` -- the web side: one short-timeout request. + +See docs/IPC_CONTROL_SOCKET.md for the protocol, the security model and the +stage plan. +""" diff --git a/src/ipc/client.py b/src/ipc/client.py new file mode 100644 index 00000000..a38d48fb --- /dev/null +++ b/src/ipc/client.py @@ -0,0 +1,200 @@ +"""The web side of the control socket: one request, a short timeout, no retries. + +Every failure -- no socket (the display is stopped, or predates the socket), +a refused or timed-out connection, a reply that breaks the contract, or an +error the display returned -- raises :class:`ControlError` with a short +``reason``, and the caller falls back to the file mailbox. Nothing here +blocks for longer than ``timeout`` in total. +""" + +from __future__ import annotations + +import socket +import time +import uuid +from typing import Any, Dict, List, Mapping, Optional, Sequence + +from src.ipc.contract import ( + MAX_MESSAGE_BYTES, + PROTOCOL_VERSION, + SUPPORTED_VERSIONS, + Command, + FrameReader, + ProtocolError, + Request, + Response, + client_socket_paths, + decode_message, + encode_message, + parse_args, + socket_supported, +) + +#: Total budget for one request: connect, send and the reply. The display +#: answers from a thread that does no rendering, normally within a few +#: milliseconds; this only bounds a wedged one. The web route then falls back +#: to the mailbox, so a timeout costs this much latency and nothing else. +DEFAULT_TIMEOUT_SECONDS = 1.0 + + +class ControlError(Exception): + """The socket could not carry the request. ``reason`` is a short code. + + Transport reasons: ``disabled``, ``unsupported``, ``no_socket``, + ``refused``, ``timeout``, ``closed``, ``bad_response``, ``invalid_request``. + When the display answered with an error, ``reason`` is that error's + :class:`~src.ipc.contract.ErrorCode` (``busy``, ``unknown_command``, ...). + """ + + def __init__(self, reason: str, message: str = ''): + super().__init__(reason, message) + self.reason = reason + self.message = message + + def __str__(self) -> str: + return f'{self.reason}: {self.message}' if self.message else self.reason + + +def request(cmd: str, args: Optional[Mapping[str, Any]] = None, *, + request_id: Optional[str] = None, + timeout: float = DEFAULT_TIMEOUT_SECONDS, + paths: Optional[Sequence[str]] = None) -> Dict[str, Any]: + """Send one command and return its ``result``. Raises :class:`ControlError`.""" + args = dict(args or {}) + request_id = request_id or str(uuid.uuid4()) + try: + # Refuse locally what the display would refuse: a malformed id + # (callers may pass their own) or arguments that break the contract. + envelope = Request.from_dict({'v': PROTOCOL_VERSION, 'id': request_id, + 'cmd': cmd, 'args': args}) + parse_args(cmd, args) + payload = encode_message(envelope.to_dict()) + except ProtocolError as e: + raise ControlError('invalid_request', e.message) from None + + if not socket_supported(): + raise ControlError('unsupported', 'no Unix sockets on this platform') + candidates: List[str] = list(paths) if paths is not None else client_socket_paths() + if not candidates: + raise ControlError('disabled', 'the control socket is turned off') + + deadline = time.monotonic() + timeout + sock = _connect(candidates, deadline) + try: + response = _exchange(sock, payload, deadline) + finally: + sock.close() + + # A refusal before the request was read (forbidden, too many + # connections) carries no id. + if response.id != request_id and not (response.id is None and not response.ok): + raise ControlError('bad_response', 'the reply is for a different request') + if not response.ok: + error = response.error + raise ControlError(error.code if error else 'bad_response', + error.message if error else '') + return dict(response.result or {}) + + +def _remaining(deadline: float) -> float: + left = deadline - time.monotonic() + if left <= 0: + raise ControlError('timeout', 'no reply in time') + return left + + +def _connect(paths: Sequence[str], deadline: float) -> socket.socket: + last = ControlError('no_socket', 'the display is not serving the control socket') + for path in paths: + sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + try: + sock.settimeout(_remaining(deadline)) + sock.connect(path) + return sock + except (FileNotFoundError, NotADirectoryError): + sock.close() + continue + except ConnectionRefusedError: + sock.close() + last = ControlError('refused', f'nothing is listening at {path}') + except BlockingIOError: + # EAGAIN: the listen backlog is full -- a live but swamped display. + sock.close() + raise ControlError('busy', 'the display is not accepting connections') from None + except PermissionError: + sock.close() + last = ControlError('refused', f'no permission to connect to {path}') + except socket.timeout: + sock.close() + raise ControlError('timeout', 'connect timed out') from None + except ControlError: + sock.close() + raise + except OSError as e: + sock.close() + last = ControlError('refused', f'{path}: {e}') + raise last + + +def _exchange(sock: socket.socket, payload: bytes, deadline: float) -> Response: + try: + sock.settimeout(_remaining(deadline)) + sock.sendall(payload) + reader = FrameReader(MAX_MESSAGE_BYTES) + while True: + sock.settimeout(_remaining(deadline)) + data = sock.recv(4096) + if not data: + raise ControlError('closed', 'the display closed the connection') + lines = reader.feed(data) + if lines: + return Response.from_dict(decode_message(lines[0])) + except socket.timeout: + raise ControlError('timeout', 'no reply in time') from None + except ProtocolError as e: + raise ControlError('bad_response', e.message) from None + except ControlError: + raise + except OSError as e: + raise ControlError('closed', str(e)) from None + + +# -- commands --------------------------------------------------------------------------- + +def on_demand_start(request_id: str, plugin_id: Optional[str], mode: Optional[str], + duration: Any = None, pinned: bool = False, *, + timeout: float = DEFAULT_TIMEOUT_SECONDS, + paths: Optional[Sequence[str]] = None) -> Dict[str, Any]: + """Ask the display to show a plugin now. Returns the ack; raises :class:`ControlError`. + + ``request_id`` doubles as the on-demand request id, so a request that a + timed-out caller then also writes to the mailbox is processed only once. + """ + args = {'plugin_id': plugin_id, 'mode': mode, 'duration': duration, 'pinned': pinned} + return request(Command.ON_DEMAND_START, args, request_id=request_id, + timeout=timeout, paths=paths) + + +def on_demand_stop(request_id: str, *, timeout: float = DEFAULT_TIMEOUT_SECONDS, + paths: Optional[Sequence[str]] = None) -> Dict[str, Any]: + """Ask the display to end on-demand. Returns the ack; raises :class:`ControlError`.""" + return request(Command.ON_DEMAND_STOP, {}, request_id=request_id, + timeout=timeout, paths=paths) + + +def on_demand_status(*, timeout: float = DEFAULT_TIMEOUT_SECONDS, + paths: Optional[Sequence[str]] = None) -> Dict[str, Any]: + """The display's live on-demand state. Raises :class:`ControlError`.""" + return request(Command.ON_DEMAND_STATUS, {}, timeout=timeout, paths=paths) + + +def ping(*, timeout: float = DEFAULT_TIMEOUT_SECONDS, + paths: Optional[Sequence[str]] = None) -> Dict[str, Any]: + return request(Command.PING, {}, timeout=timeout, paths=paths) + + +def hello(client: str = 'web', *, timeout: float = DEFAULT_TIMEOUT_SECONDS, + paths: Optional[Sequence[str]] = None) -> Dict[str, Any]: + """Version negotiation: the result's ``version`` is the one both sides speak.""" + return request(Command.HELLO, {'versions': list(SUPPORTED_VERSIONS), 'client': client}, + timeout=timeout, paths=paths) diff --git a/src/ipc/contract.py b/src/ipc/contract.py new file mode 100644 index 00000000..3b1d4ad8 --- /dev/null +++ b/src/ipc/contract.py @@ -0,0 +1,527 @@ +"""The control socket's contract: versioned messages, framing and location. + +Both processes import this module -- the display serves the socket +(:mod:`src.ipc.server`) and the web interface calls it +(:mod:`src.ipc.client`) -- so it is the one definition of what goes over the +wire. Standard library only, and no import of the rest of ``src``. + +Wire format (protocol version 1) +-------------------------------- +One JSON object per line (newline-delimited JSON), UTF-8, at most +:data:`MAX_MESSAGE_BYTES` per line including the newline. Messages are +encoded with ``ensure_ascii``, so a newline never appears inside one. + +Request:: + + {"v": 1, "id": "<1-128 chars>", "cmd": "on_demand.start", "args": {...}} + +Response, always carrying the request's ``id`` (``null`` when the request +could not be parsed far enough to have one):: + + {"v": 1, "id": "...", "ok": true, "result": {...}} + {"v": 1, "id": "...", "ok": false, "error": {"code": "...", "message": "..."}} + +A connection may carry several requests; each gets exactly one response, in +order. Commands that change what the panel shows are *acknowledged*, not +completed: ``{"accepted": true, "request_id": ...}`` means the render thread +has the command queued and will apply it at its next on-demand check. Its +outcome is published the way it always was (``display_on_demand_state``, +later the state stream). + +See docs/IPC_CONTROL_SOCKET.md for the full description. +""" + +from __future__ import annotations + +import json +import math +import os +import tempfile +from dataclasses import dataclass, field +from typing import Any, Dict, List, Mapping, Optional, Tuple, TypedDict, TypeGuard, Union + + +# -- versions and limits ------------------------------------------------------- + +#: The protocol version this code speaks by default. +PROTOCOL_VERSION = 1 + +#: Every version this code can speak; ``hello`` picks the highest common one. +SUPPORTED_VERSIONS: Tuple[int, ...] = (1,) + +#: The largest message either side sends or accepts, newline included. A +#: stage-1 message is well under 1 KiB; this only bounds a broken or hostile +#: peer, so a reader never buffers more than this per connection. +MAX_MESSAGE_BYTES = 64 * 1024 + +#: Longest request id. Ids are also the on-demand ``request_id``, which the +#: display logs and stores, so they are kept short. +MAX_ID_LENGTH = 128 + +#: Longest plugin id or mode name an on-demand command may carry. +MAX_NAME_LENGTH = 128 + + +# -- where the socket lives ------------------------------------------------------ + +#: ``RuntimeDirectory=ledmatrix`` in ledmatrix.service creates this (tmpfs, +#: root-owned, 0755); a display under an older unit creates it itself, as it +#: does for the heartbeat (src/display_watchdog.py). +DEFAULT_SOCKET_DIR = '/run/ledmatrix' +SOCKET_NAME = 'control.sock' +DEFAULT_SOCKET_PATH = DEFAULT_SOCKET_DIR + '/' + SOCKET_NAME + +#: Overrides the socket path for both processes (a dev checkout, a second +#: instance, tests). One of :data:`DISABLED_VALUES` turns the socket off: the +#: display does not serve it and the web interface goes straight to the +#: file mailbox. +SOCKET_PATH_ENV = 'LEDMATRIX_CONTROL_SOCKET' +DISABLED_VALUES = frozenset({'off', '0', 'false', 'no', 'none', 'disabled'}) + + +def socket_supported() -> bool: + """Whether this platform has Unix sockets at all (Windows Python does not).""" + import socket + return os.name == 'posix' and hasattr(socket, 'AF_UNIX') + + +def socket_disabled(environ: Optional[Mapping[str, str]] = None) -> bool: + """True when :data:`SOCKET_PATH_ENV` switches the socket off.""" + env = os.environ if environ is None else environ + value = (env.get(SOCKET_PATH_ENV) or '').strip() + return value.lower() in DISABLED_VALUES + + +def configured_socket_path(environ: Optional[Mapping[str, str]] = None) -> Optional[str]: + """The path :data:`SOCKET_PATH_ENV` names, or None when it is unset or 'off'.""" + env = os.environ if environ is None else environ + value = (env.get(SOCKET_PATH_ENV) or '').strip() + if not value or value.lower() in DISABLED_VALUES: + return None + return value + + +def dev_socket_path(uid: Optional[int] = None) -> str: + """Where a display that cannot use /run/ledmatrix serves the socket. + + A per-user directory under the temp dir, so a dev checkout run as an + ordinary user (``python3 run.py -e``) and its web interface, run by the + same user, find each other with no configuration. + """ + if uid is None: + getuid = getattr(os, 'getuid', None) + uid = getuid() if getuid is not None else 0 + return os.path.join(tempfile.gettempdir(), f'ledmatrix-{uid}', SOCKET_NAME) + + +def client_socket_paths(environ: Optional[Mapping[str, str]] = None) -> List[str]: + """The paths a client tries, in order; empty when the socket is off.""" + if socket_disabled(environ): + return [] + configured = configured_socket_path(environ) + if configured: + return [configured] + return [DEFAULT_SOCKET_PATH, dev_socket_path()] + + +# -- commands and error codes ---------------------------------------------------- + +class Command: + """Command names. Dotted names group a feature's commands.""" + HELLO = 'hello' + PING = 'ping' + ON_DEMAND_START = 'on_demand.start' + ON_DEMAND_STOP = 'on_demand.stop' + ON_DEMAND_STATUS = 'on_demand.status' + + +#: Every command version 1 defines, in the order ``hello`` reports them. +COMMANDS: Tuple[str, ...] = ( + Command.HELLO, + Command.PING, + Command.ON_DEMAND_START, + Command.ON_DEMAND_STOP, + Command.ON_DEMAND_STATUS, +) + +#: Commands that are queued for the render thread and answered with an ack. +QUEUED_COMMANDS = frozenset({Command.ON_DEMAND_START, Command.ON_DEMAND_STOP}) + + +class ErrorCode: + """``error.code`` values. Clients branch on these, never on the message.""" + BAD_JSON = 'bad_json' # a line that is not a JSON object + BAD_REQUEST = 'bad_request' # the envelope is malformed + MESSAGE_TOO_LARGE = 'message_too_large' # over MAX_MESSAGE_BYTES + UNSUPPORTED_VERSION = 'unsupported_version' # no version in common + UNKNOWN_COMMAND = 'unknown_command' + INVALID_ARGS = 'invalid_args' + BUSY = 'busy' # queue full / too many clients + FORBIDDEN = 'forbidden' # peer credentials refused + INTERNAL = 'internal' # a bug on the display side + + +class ProtocolError(Exception): + """A message that breaks the contract. ``code`` is an :class:`ErrorCode`.""" + + def __init__(self, code: str, message: str, request_id: Optional[str] = None): + super().__init__(code, message, request_id) + self.code = code + self.message = message + self.request_id = request_id + + def __str__(self) -> str: + return f'{self.code}: {self.message}' + + +# -- the envelope ------------------------------------------------------------------ + +def _is_int(value: Any) -> TypeGuard[int]: + return isinstance(value, int) and not isinstance(value, bool) + + +def _valid_id(value: Any) -> bool: + return (isinstance(value, str) and 0 < len(value) <= MAX_ID_LENGTH + and value.isprintable()) + + +@dataclass(frozen=True) +class Request: + """``{v, id, cmd, args}``.""" + id: str + cmd: str + args: Dict[str, Any] = field(default_factory=dict) + v: int = PROTOCOL_VERSION + + def to_dict(self) -> Dict[str, Any]: + return {'v': self.v, 'id': self.id, 'cmd': self.cmd, 'args': dict(self.args)} + + @classmethod + def from_dict(cls, obj: Any) -> 'Request': + """Validate an envelope. Raises :class:`ProtocolError`. + + The version is checked by the server, not here, so that ``hello`` + can negotiate across versions. + """ + if not isinstance(obj, dict): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'a request must be a JSON object') + raw_id = obj.get('id') + request_id = raw_id if _valid_id(raw_id) else None + if request_id is None: + raise ProtocolError(ErrorCode.BAD_REQUEST, + f'id must be a printable string of 1-{MAX_ID_LENGTH} characters') + version = obj.get('v') + if not _is_int(version): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'v must be an integer', request_id) + cmd = obj.get('cmd') + if not isinstance(cmd, str) or not cmd: + raise ProtocolError(ErrorCode.BAD_REQUEST, 'cmd must be a non-empty string', request_id) + args = obj.get('args', {}) + if args is None: + args = {} + if not isinstance(args, dict): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'args must be a JSON object', request_id) + return cls(id=request_id, cmd=cmd, args=args, v=version) + + +@dataclass(frozen=True) +class ErrorInfo: + code: str + message: str + + def to_dict(self) -> Dict[str, str]: + return {'code': self.code, 'message': self.message} + + +@dataclass(frozen=True) +class Response: + """``{v, id, ok, result}`` or ``{v, id, ok: false, error: {code, message}}``.""" + id: Optional[str] + ok: bool + result: Optional[Dict[str, Any]] = None + error: Optional[ErrorInfo] = None + v: int = PROTOCOL_VERSION + + @classmethod + def success(cls, request_id: Optional[str], result: Mapping[str, Any], + v: int = PROTOCOL_VERSION) -> 'Response': + return cls(id=request_id, ok=True, result=dict(result), v=v) + + @classmethod + def failure(cls, request_id: Optional[str], code: str, message: str, + v: int = PROTOCOL_VERSION) -> 'Response': + return cls(id=request_id, ok=False, error=ErrorInfo(code, message), v=v) + + def to_dict(self) -> Dict[str, Any]: + out: Dict[str, Any] = {'v': self.v, 'id': self.id, 'ok': self.ok} + if self.ok: + out['result'] = dict(self.result or {}) + else: + error = self.error or ErrorInfo(ErrorCode.INTERNAL, 'unknown error') + out['error'] = error.to_dict() + return out + + @classmethod + def from_dict(cls, obj: Any) -> 'Response': + """Validate a response. Raises :class:`ProtocolError` (BAD_REQUEST).""" + if not isinstance(obj, dict): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'a response must be a JSON object') + version = obj.get('v') + if not _is_int(version): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'v must be an integer') + raw_id = obj.get('id') + if raw_id is not None and not isinstance(raw_id, str): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'id must be a string or null') + ok = obj.get('ok') + if not isinstance(ok, bool): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'ok must be a boolean') + if ok: + result = obj.get('result', {}) + if not isinstance(result, dict): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'result must be a JSON object') + return cls(id=raw_id, ok=True, result=result, v=version) + error = obj.get('error') + if (not isinstance(error, dict) or not isinstance(error.get('code'), str) + or not isinstance(error.get('message', ''), str)): + raise ProtocolError(ErrorCode.BAD_REQUEST, 'error must be {code, message}') + return cls(id=raw_id, ok=False, + error=ErrorInfo(error['code'], error.get('message', '')), v=version) + + +# -- command arguments ------------------------------------------------------------- + +def _optional_name(args: Mapping[str, Any], key: str) -> Optional[str]: + value = args.get(key) + if value is None or value == '': + return None + if not isinstance(value, str) or len(value) > MAX_NAME_LENGTH or not value.isprintable(): + raise ProtocolError(ErrorCode.INVALID_ARGS, + f'{key} must be a printable string of at most ' + f'{MAX_NAME_LENGTH} characters') + return value + + +def _optional_duration(value: Any) -> Optional[float]: + """Seconds, or None for "until stopped". 0 means the same as None. + + Numbers and numeric strings are accepted, the same as the REST route and + the file mailbox take them; anything else is refused rather than guessed. + """ + if value is None or value == '': + return None + if isinstance(value, bool): + raise ProtocolError(ErrorCode.INVALID_ARGS, 'duration must be a number of seconds') + try: + seconds = float(value) + except (TypeError, ValueError): + raise ProtocolError(ErrorCode.INVALID_ARGS, + 'duration must be a number of seconds') from None + if not math.isfinite(seconds) or seconds < 0: + raise ProtocolError(ErrorCode.INVALID_ARGS, + 'duration must be a finite, non-negative number of seconds') + return seconds or None + + +@dataclass(frozen=True) +class HelloArgs: + """``hello``: the versions the client speaks, and a name for the logs.""" + versions: Tuple[int, ...] = (PROTOCOL_VERSION,) + client: str = '' + + def to_dict(self) -> Dict[str, Any]: + return {'versions': list(self.versions), 'client': self.client} + + @classmethod + def from_dict(cls, args: Mapping[str, Any]) -> 'HelloArgs': + versions = args.get('versions', [PROTOCOL_VERSION]) + if (not isinstance(versions, list) or not versions or len(versions) > 32 + or not all(_is_int(v) for v in versions)): + raise ProtocolError(ErrorCode.INVALID_ARGS, 'versions must be a list of integers') + client = args.get('client', '') + if not isinstance(client, str) or len(client) > MAX_NAME_LENGTH: + raise ProtocolError(ErrorCode.INVALID_ARGS, 'client must be a short string') + return cls(versions=tuple(versions), client=client) + + +@dataclass(frozen=True) +class OnDemandStartArgs: + """``on_demand.start``: show a plugin (or one of its modes) now. + + The same fields the file mailbox carries. At least one of ``plugin_id`` + and ``mode`` is required; the display resolves the other. + """ + plugin_id: Optional[str] = None + mode: Optional[str] = None + duration: Optional[float] = None + pinned: bool = False + + def to_dict(self) -> Dict[str, Any]: + return {'plugin_id': self.plugin_id, 'mode': self.mode, + 'duration': self.duration, 'pinned': self.pinned} + + @classmethod + def from_dict(cls, args: Mapping[str, Any]) -> 'OnDemandStartArgs': + plugin_id = _optional_name(args, 'plugin_id') + mode = _optional_name(args, 'mode') + if plugin_id is None and mode is None: + raise ProtocolError(ErrorCode.INVALID_ARGS, 'plugin_id or mode is required') + pinned = args.get('pinned', False) + if pinned is None: + pinned = False + if not isinstance(pinned, bool): + raise ProtocolError(ErrorCode.INVALID_ARGS, 'pinned must be a boolean') + return cls(plugin_id=plugin_id, mode=mode, + duration=_optional_duration(args.get('duration')), pinned=pinned) + + +@dataclass(frozen=True) +class OnDemandStopArgs: + """``on_demand.stop``: end the on-demand session and resume rotation.""" + + def to_dict(self) -> Dict[str, Any]: + return {} + + @classmethod + def from_dict(cls, args: Mapping[str, Any]) -> 'OnDemandStopArgs': + return cls() + + +@dataclass(frozen=True) +class NoArgs: + """``ping`` and ``on_demand.status`` take no arguments (extra ones are ignored).""" + + def to_dict(self) -> Dict[str, Any]: + return {} + + @classmethod + def from_dict(cls, args: Mapping[str, Any]) -> 'NoArgs': + return cls() + + +CommandArgs = Union[HelloArgs, OnDemandStartArgs, OnDemandStopArgs, NoArgs] + +_ARG_TYPES: Dict[str, Any] = { + Command.HELLO: HelloArgs, + Command.PING: NoArgs, + Command.ON_DEMAND_START: OnDemandStartArgs, + Command.ON_DEMAND_STOP: OnDemandStopArgs, + Command.ON_DEMAND_STATUS: NoArgs, +} + + +def parse_args(cmd: str, args: Mapping[str, Any]) -> CommandArgs: + """Typed arguments for ``cmd``. Raises :class:`ProtocolError`.""" + arg_type = _ARG_TYPES.get(cmd) + if arg_type is None: + raise ProtocolError(ErrorCode.UNKNOWN_COMMAND, f'unknown command: {cmd[:64]}') + parsed: CommandArgs = arg_type.from_dict(args) + return parsed + + +def on_demand_request(request_id: str, args: Union[OnDemandStartArgs, OnDemandStopArgs], + timestamp: float) -> Dict[str, Any]: + """The file-mailbox payload for a queued on-demand command. + + The display hands socket commands to the same code that handles the + mailbox (``DisplayController._handle_on_demand_request``), so a command + behaves identically whichever way it arrived, and a request that came + both ways (a client that timed out and fell back) is processed once: the + request id is the same. + """ + if isinstance(args, OnDemandStartArgs): + return {'request_id': request_id, 'action': 'start', 'plugin_id': args.plugin_id, + 'mode': args.mode, 'duration': args.duration, 'pinned': args.pinned, + 'timestamp': timestamp, 'source': 'socket'} + return {'request_id': request_id, 'action': 'stop', 'timestamp': timestamp, + 'source': 'socket'} + + +# -- results ----------------------------------------------------------------------- + +class HelloResult(TypedDict): + version: int + versions: List[int] + commands: List[str] + max_message_bytes: int + server: str + + +class PingResult(TypedDict): + pong: bool + + +class AckResult(TypedDict): + """The answer to a queued command: the render thread will apply it.""" + accepted: bool + request_id: str + queued: int + + +def negotiate_version(client_versions: Tuple[int, ...]) -> Optional[int]: + """The highest version both sides speak, or None.""" + common = set(client_versions) & set(SUPPORTED_VERSIONS) + return max(common) if common else None + + +# -- framing ----------------------------------------------------------------------- + +def encode_message(obj: Mapping[str, Any]) -> bytes: + """One newline-terminated JSON line. Raises :class:`ProtocolError` when too big.""" + try: + text = json.dumps(obj, separators=(',', ':'), ensure_ascii=True, allow_nan=False) + except (TypeError, ValueError) as e: + raise ProtocolError(ErrorCode.BAD_REQUEST, f'message is not JSON-serialisable: {e}') from None + data = text.encode('ascii') + b'\n' + if len(data) > MAX_MESSAGE_BYTES: + raise ProtocolError(ErrorCode.MESSAGE_TOO_LARGE, + f'message is {len(data)} bytes; the limit is {MAX_MESSAGE_BYTES}') + return data + + +def decode_message(line: bytes) -> Dict[str, Any]: + """Parse one line (newline optional). Raises :class:`ProtocolError` (BAD_JSON).""" + try: + obj = json.loads(line.decode('utf-8')) + except ValueError: # UnicodeDecodeError and JSONDecodeError are both ValueErrors + raise ProtocolError(ErrorCode.BAD_JSON, 'not valid UTF-8 JSON') from None + if not isinstance(obj, dict): + raise ProtocolError(ErrorCode.BAD_JSON, 'a message must be a JSON object') + return obj + + +class FrameReader: + """Splits a byte stream into lines, never holding more than one message. + + ``feed()`` returns the complete lines (without their newlines) the new + bytes finished, and raises :class:`ProtocolError` (MESSAGE_TOO_LARGE) as + soon as a line is longer than the limit, newline or not, so a peer that + never sends one cannot make the reader buffer without bound. + """ + + def __init__(self, max_bytes: int = MAX_MESSAGE_BYTES): + self._max = max_bytes + self._buffer = bytearray() + + @property + def pending(self) -> int: + """Bytes of an unfinished message held.""" + return len(self._buffer) + + def feed(self, data: bytes) -> List[bytes]: + self._buffer.extend(data) + lines: List[bytes] = [] + while True: + newline = self._buffer.find(b'\n') + if newline < 0: + break + if newline + 1 > self._max: + raise ProtocolError(ErrorCode.MESSAGE_TOO_LARGE, + f'message exceeds {self._max} bytes') + line = bytes(self._buffer[:newline]) + del self._buffer[:newline + 1] + if line.strip(): + lines.append(line) + if len(self._buffer) >= self._max: + raise ProtocolError(ErrorCode.MESSAGE_TOO_LARGE, + f'message exceeds {self._max} bytes') + return lines diff --git a/src/ipc/server.py b/src/ipc/server.py new file mode 100644 index 00000000..7fc08e50 --- /dev/null +++ b/src/ipc/server.py @@ -0,0 +1,643 @@ +"""The display side of the control socket. + +A small threaded server on a Unix stream socket (``/run/ledmatrix/control.sock`` +by default; see :mod:`src.ipc.contract` for the protocol). It never touches +rendering: a command that changes the panel is validated, put on a bounded +queue and acknowledged, and the render thread drains that queue at the point +where it reads the file mailbox (``DisplayController._poll_on_demand_requests``), +handing each command to the same code. Queries (``on_demand.status``) are +answered from a snapshot callable the display provides. + +Robustness rules, because this runs inside the display process: + +* every connection has its own daemon thread, at most :data:`MAX_CLIENTS` at + once; one more is told ``busy`` and closed; +* every read and write has a timeout, a message must arrive whole within + :data:`MESSAGE_TIMEOUT_SECONDS`, and an idle connection is closed after + :data:`IDLE_TIMEOUT_SECONDS` -- a slow or stuck client costs one thread for + a few seconds, never the render loop; +* a line that is not JSON is answered with ``bad_json`` and the connection + carries on; a line over the size limit closes the connection; a client + that disconnects mid-message is simply dropped; +* no exception from a handler leaves the connection thread. + +Who may connect (see docs/IPC_CONTROL_SOCKET.md, "Security model"): the +socket file is ``0660`` and group-owned by the group the display and the web +interface share -- the cache directory's group, the same rule DiskCache uses +for the files it shares -- so the kernel refuses everyone else at connect(). +Where the kernel reports the peer's credentials (``SO_PEERCRED``, Linux) the +server checks them again: root, its own user, or a member of that group. +""" + +from __future__ import annotations + +import logging +import os +import queue +import socket +import stat +import struct +import threading +import time +from dataclasses import dataclass +from typing import Any, Callable, Dict, FrozenSet, List, Mapping, Optional, Union + +from src.ipc.contract import ( + COMMANDS, + DEFAULT_SOCKET_DIR, + DEFAULT_SOCKET_PATH, + MAX_MESSAGE_BYTES, + PROTOCOL_VERSION, + QUEUED_COMMANDS, + SUPPORTED_VERSIONS, + AckResult, + Command, + ErrorCode, + FrameReader, + HelloArgs, + HelloResult, + OnDemandStartArgs, + OnDemandStopArgs, + ProtocolError, + Request, + Response, + configured_socket_path, + decode_message, + dev_socket_path, + encode_message, + negotiate_version, + on_demand_request, + parse_args, + socket_disabled, + socket_supported, +) + +logger = logging.getLogger(__name__) + +#: Concurrent connections served. The web interface opens one per request +#: and closes it; this only bounds a misbehaving client. +MAX_CLIENTS = 8 + +#: Commands waiting for the render thread. It drains them at least every +#: 0.25 s, so a full queue means the render thread is stuck, and the client +#: is told ``busy`` (and falls back to the mailbox) instead of piling up work. +QUEUE_SIZE = 16 + +#: Timeout for one recv()/send() on a connection. +IO_TIMEOUT_SECONDS = 2.0 + +#: A message must arrive whole within this long of its first byte. +MESSAGE_TIMEOUT_SECONDS = 5.0 + +#: A connection with no message in progress is closed after this long. +IDLE_TIMEOUT_SECONDS = 10.0 + +#: How often the accept loop wakes to notice close(). +_ACCEPT_POLL_SECONDS = 0.5 + +_LISTEN_BACKLOG = 64 + + +# -- queued work --------------------------------------------------------------------- + +@dataclass(frozen=True) +class QueuedCommand: + """A command waiting for the render thread.""" + request_id: str + cmd: str + args: Union[OnDemandStartArgs, OnDemandStopArgs] + received_at: float # time.time() when it was accepted + peer_uid: Optional[int] = None + + def as_on_demand_request(self) -> Dict[str, Any]: + """The mailbox-shaped payload the display's on-demand handler takes.""" + return on_demand_request(self.request_id, self.args, self.received_at) + + +# -- peer credentials ------------------------------------------------------------------ + +@dataclass(frozen=True) +class PeerCredentials: + pid: int + uid: int + gid: int + + +def peer_credentials(conn: socket.socket) -> Optional[PeerCredentials]: + """The connecting process's pid/uid/gid, where the kernel reports them. + + ``SO_PEERCRED`` is Linux's; elsewhere this is None and the socket file's + mode is the only gate. + """ + option = getattr(socket, 'SO_PEERCRED', None) + if option is None: + return None + try: + raw = conn.getsockopt(socket.SOL_SOCKET, option, struct.calcsize('3i')) + pid, uid, gid = struct.unpack('3i', raw) + except (OSError, struct.error): + return None + return PeerCredentials(pid=pid, uid=uid, gid=gid) + + +def process_groups(pid: int) -> Optional[FrozenSet[int]]: + """A process's supplementary groups, from /proc; None when unreadable. + + The web service's primary group is normally its user's own; the shared + group is a supplementary one, which ``SO_PEERCRED`` does not report. + """ + try: + with open(f'/proc/{int(pid)}/status', 'r', encoding='ascii', errors='replace') as f: + for line in f: + if line.startswith('Groups:'): + return frozenset(int(g) for g in line.split()[1:] if g.isdigit()) + except (OSError, ValueError): + return None + return frozenset() + + +def user_in_group(uid: int, gid: int) -> bool: + """Whether the account ``uid`` is listed in group ``gid`` (the group database).""" + try: + import grp + import pwd + name = pwd.getpwuid(uid).pw_name + group = grp.getgrgid(gid) + except (ImportError, KeyError, OSError): + return False + return name in group.gr_mem or pwd.getpwuid(uid).pw_gid == gid + + +def peer_allowed(cred: PeerCredentials, own_uid: int, allowed_gid: Optional[int], + groups: Optional[FrozenSet[int]] = None, + in_group: Callable[[int, int], bool] = user_in_group) -> bool: + """The permission model: root, the server's own user, or the shared group. + + ``groups`` are the peer's supplementary groups (from /proc); when they + could not be read the group database decides instead. + """ + if cred.uid == 0 or cred.uid == own_uid: + return True + if allowed_gid is None: + return False + if cred.gid == allowed_gid: + return True + if groups is not None: + return allowed_gid in groups + return in_group(cred.uid, allowed_gid) + + +def resolve_socket_group(cache_dir: Optional[str]) -> Optional[int]: + """The group the socket should belong to: the one the two services share. + + The cache directory's group when the directory is group-writable -- + the rule DiskCache applies to every file the display shares with the web + interface (``root:ledmatrix 2775`` on an installed device). Otherwise the + project directory's group (``get_shared_group_gid``), which config files + use. None when neither is known: then only root and the display's own + user can connect. + """ + if cache_dir: + try: + st = os.stat(cache_dir) + if st.st_mode & stat.S_IWGRP: + return st.st_gid + except OSError: + pass + try: + from src.common.permission_utils import get_shared_group_gid + return get_shared_group_gid() + except ImportError: # pragma: no cover - src is always importable here + return None + + +def server_socket_path(environ: Optional[Mapping[str, str]] = None) -> Optional[str]: + """Where the display should serve the socket; None when it should not. + + :data:`~src.ipc.contract.SOCKET_PATH_ENV` wins. Otherwise + /run/ledmatrix/control.sock when the display can create or write that + directory (root, which an installed display always is), and the per-user + dev path otherwise (an emulator run from a checkout). + """ + if not socket_supported() or socket_disabled(environ): + return None + configured = configured_socket_path(environ) + if configured: + return configured + geteuid = getattr(os, 'geteuid', None) + if (geteuid is not None and geteuid() == 0) or os.access(DEFAULT_SOCKET_DIR, os.W_OK): + return DEFAULT_SOCKET_PATH + return dev_socket_path() + + +# -- the server ------------------------------------------------------------------------ + +StatusProvider = Callable[[], Dict[str, Any]] + + +class ControlServer: + """Serves the control socket on background threads. + + ``start()`` binds and starts accepting; ``drain()`` (render thread) takes + the queued commands; ``close()`` stops and removes the socket file. + """ + + def __init__(self, path: str, status_provider: Optional[StatusProvider] = None, + group: Optional[int] = None, *, queue_size: int = QUEUE_SIZE, + max_clients: int = MAX_CLIENTS, io_timeout: float = IO_TIMEOUT_SECONDS, + message_timeout: float = MESSAGE_TIMEOUT_SECONDS, + idle_timeout: float = IDLE_TIMEOUT_SECONDS, + check_peer: bool = True): + self.path = path + self._status_provider = status_provider + self._group = group + self._queue: 'queue.Queue[QueuedCommand]' = queue.Queue(maxsize=queue_size) + self._pending = threading.Event() + self._slots = threading.BoundedSemaphore(max_clients) + self._io_timeout = io_timeout + self._message_timeout = message_timeout + self._idle_timeout = idle_timeout + self._check_peer = check_peer + self._sock: Optional[socket.socket] = None + self._identity: Optional[tuple] = None # (st_dev, st_ino) of our socket file + self._thread: Optional[threading.Thread] = None + self._stopping = threading.Event() + self._own_uid = os.geteuid() if hasattr(os, 'geteuid') else -1 + + # -- lifecycle ------------------------------------------------------------- + + @property + def running(self) -> bool: + return self._thread is not None and self._thread.is_alive() + + @property + def socket_mode(self) -> int: + """0660 with a shared group; 0600 (the display's user only) without one.""" + return 0o660 if self._group is not None else 0o600 + + def start(self) -> bool: + """Bind and start serving. False (logged) when the socket cannot be served. + + Never raises: without the socket the web interface uses the file + mailbox, exactly as before. + """ + if not socket_supported(): + logger.debug("Control socket not started: no Unix sockets on this platform") + return False + try: + self._prepare_directory() + if not self._clear_stale_socket(): + return False + self._bind() + except OSError as e: + logger.warning("Control socket not started at %s (%s); the web interface " + "will use the file mailbox", self.path, e) + self._close_socket() + return False + self._stopping.clear() + self._thread = threading.Thread(target=self._accept_loop, name='ledmatrix-ipc', + daemon=True) + self._thread.start() + logger.info("Control socket listening at %s (mode %o, group %s)", + self.path, self.socket_mode, + self._group if self._group is not None else 'none') + return True + + def close(self) -> None: + """Stop accepting and remove the socket file (only if it is still ours).""" + self._stopping.set() + self._close_socket() + thread = self._thread + if thread is not None and thread is not threading.current_thread(): + thread.join(timeout=2.0) + self._thread = None + if self._identity is not None: + try: + st = os.lstat(self.path) + if (st.st_dev, st.st_ino) == self._identity: + os.unlink(self.path) + except OSError: + pass + self._identity = None + + def _close_socket(self) -> None: + sock, self._sock = self._sock, None + if sock is not None: + try: + sock.close() + except OSError: + pass + + def _prepare_directory(self) -> None: + directory = os.path.dirname(os.path.abspath(self.path)) + if self.path == dev_socket_path(): + # The dev path is in the shared temp dir: private to this user, + # and refused if someone else got there first. + os.makedirs(directory, mode=0o700, exist_ok=True) + self._check_private_directory(directory) + elif not os.path.isdir(directory): + # /run/ledmatrix under a unit that predates RuntimeDirectory= (the + # display is root and makes it, as it does for the heartbeat), or + # a configured path. 0755: the web interface only needs to reach + # the socket; the socket's own mode decides who may connect. + os.makedirs(directory, mode=0o755, exist_ok=True) + + def _check_private_directory(self, directory: str) -> None: + """Refuse a dev directory someone else made (it lives in a shared /tmp).""" + st = os.lstat(directory) + if stat.S_ISLNK(st.st_mode) or not stat.S_ISDIR(st.st_mode): + raise OSError(f'{directory} is not a plain directory') + if hasattr(os, 'geteuid') and st.st_uid != os.geteuid(): + raise OSError(f'{directory} belongs to uid {st.st_uid}, not this user') + + def _clear_stale_socket(self) -> bool: + """Remove a socket left by a display that died; never a live or foreign file.""" + try: + st = os.lstat(self.path) + except FileNotFoundError: + return True + if not stat.S_ISSOCK(st.st_mode): + logger.error("Control socket not started: %s exists and is not a socket", self.path) + return False + probe = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + probe.settimeout(0.5) + try: + probe.connect(self.path) + except OSError: + os.unlink(self.path) # nothing listening: a previous display's leftover + return True + finally: + probe.close() + logger.warning("Control socket not started: another process is serving %s", self.path) + return False + + def _bind(self) -> None: + """Bind under a temporary name, set mode and group, then rename into place. + + The rename makes the socket appear with its final permissions, never + briefly with the process umask's. + """ + tmp = f'{self.path}.{os.getpid()}.tmp' + try: + os.unlink(tmp) + except FileNotFoundError: + pass + sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + self._sock = sock + try: + sock.bind(tmp) + os.chmod(tmp, self.socket_mode) + if self._group is not None and hasattr(os, 'chown'): + try: + os.chown(tmp, -1, self._group) + except OSError as e: + # Not root and not in the group (a dev run): only this user + # (and root) can connect, which is what a dev run needs. + logger.debug("Could not give the control socket group %s: %s", + self._group, e) + # The backlog is only the kernel's queue in front of accept(); + # MAX_CLIENTS still bounds what is served. A short one makes a + # burst of clients fail connect() with EAGAIN instead of being + # answered (busy or otherwise). + sock.listen(_LISTEN_BACKLOG) + sock.settimeout(_ACCEPT_POLL_SECONDS) + os.rename(tmp, self.path) + except BaseException: + try: + os.unlink(tmp) + except OSError: + pass + raise + st = os.lstat(self.path) + self._identity = (st.st_dev, st.st_ino) + + # -- the render thread's side ------------------------------------------------ + + @property + def has_pending(self) -> bool: + """Cheap check for queued commands, for the render thread's fast path.""" + return self._pending.is_set() + + def drain(self) -> List[QueuedCommand]: + """Every queued command, oldest first. Called from the render thread.""" + commands: List[QueuedCommand] = [] + self._pending.clear() + while True: + try: + commands.append(self._queue.get_nowait()) + except queue.Empty: + break + return commands + + # -- serving ------------------------------------------------------------------- + + def _accept_loop(self) -> None: + while not self._stopping.is_set(): + sock = self._sock + if sock is None: + break + try: + conn, _ = sock.accept() + except socket.timeout: + continue + except OSError as e: + if self._stopping.is_set(): + break + logger.warning("Control socket accept failed: %s", e) + time.sleep(0.1) + continue + if not self._slots.acquire(blocking=False): + self._refuse(conn, ErrorCode.BUSY, 'too many connections') + continue + try: + threading.Thread(target=self._serve, args=(conn,), name='ledmatrix-ipc-conn', + daemon=True).start() + except RuntimeError: # can't start a thread: shed the client + self._slots.release() + self._refuse(conn, ErrorCode.BUSY, 'server overloaded') + + def _refuse(self, conn: socket.socket, code: str, message: str) -> None: + try: + conn.settimeout(0.2) + conn.sendall(encode_message(Response.failure(None, code, message).to_dict())) + except OSError: + pass + finally: + conn.close() + + def _serve(self, conn: socket.socket) -> None: + """One connection: authenticate, then answer requests until it ends.""" + try: + conn.settimeout(self._io_timeout) + peer = peer_credentials(conn) + if self._check_peer and peer is not None and not self._peer_ok(peer): + logger.warning("Control socket refused pid %d (uid %d, gid %d): not root, " + "this user or group %s", peer.pid, peer.uid, peer.gid, self._group) + self._send(conn, Response.failure(None, ErrorCode.FORBIDDEN, 'not permitted')) + return + self._read_requests(conn, peer) + except Exception: # pylint: disable=broad-except + logger.exception("Control socket connection failed") + finally: + try: + conn.close() + except OSError: + pass + self._slots.release() + + def _peer_ok(self, peer: PeerCredentials) -> bool: + groups = None + if peer.uid not in (0, self._own_uid) and self._group is not None: + groups = process_groups(peer.pid) + return peer_allowed(peer, self._own_uid, self._group, groups) + + def _read_requests(self, conn: socket.socket, peer: Optional[PeerCredentials]) -> None: + reader = FrameReader(MAX_MESSAGE_BYTES) + idle_since = time.monotonic() + message_started: Optional[float] = None + while not self._stopping.is_set(): + now = time.monotonic() + if message_started is not None and now - message_started > self._message_timeout: + logger.debug("Control socket: dropping a client too slow to send a message") + return + if message_started is None and now - idle_since > self._idle_timeout: + return + try: + data = conn.recv(4096) + except socket.timeout: + continue + except OSError: + return + if not data: + return # closed, possibly mid-message: nothing to answer + try: + lines = reader.feed(data) + except ProtocolError as e: + self._send(conn, Response.failure(None, e.code, e.message)) + return # can't find the next message boundary: hang up + for line in lines: + if not self._send(conn, self.handle_line(line, peer)): + return + if reader.pending: + if message_started is None or lines: + message_started = time.monotonic() + else: + message_started = None + idle_since = time.monotonic() + + def _send(self, conn: socket.socket, response: Response) -> bool: + try: + data = encode_message(response.to_dict()) + except ProtocolError as e: + # A status snapshot too big (or not JSON) to send is a display bug. + logger.error("Control socket response not sent: %s", e.message) + data = encode_message(Response.failure( + response.id, ErrorCode.INTERNAL, 'response could not be encoded').to_dict()) + try: + conn.sendall(data) + return True + except OSError: + return False + + # -- requests -------------------------------------------------------------------- + + def handle_line(self, line: bytes, peer: Optional[PeerCredentials] = None) -> Response: + """Answer one request line. Never raises.""" + request_id: Optional[str] = None + try: + obj = decode_message(line) + raw_id = obj.get('id') + request_id = raw_id if isinstance(raw_id, str) and len(raw_id) <= 128 else None + request = Request.from_dict(obj) + request_id = request.id + return self._dispatch(request, peer) + except ProtocolError as e: + return Response.failure(e.request_id or request_id, e.code, e.message) + except Exception: # pylint: disable=broad-except + logger.exception("Control socket handler failed") + return Response.failure(request_id, ErrorCode.INTERNAL, 'internal error') + + def _dispatch(self, request: Request, peer: Optional[PeerCredentials]) -> Response: + if request.cmd == Command.HELLO: + # Exempt from the envelope version check: this is how a client + # that speaks other versions finds out which ones we share. + hello = HelloArgs.from_dict(request.args) + version = negotiate_version(hello.versions) + if version is None: + return Response.failure( + request.id, ErrorCode.UNSUPPORTED_VERSION, + f'no common protocol version; this display speaks {list(SUPPORTED_VERSIONS)}') + result: HelloResult = { + 'version': version, + 'versions': list(SUPPORTED_VERSIONS), + 'commands': list(COMMANDS), + 'max_message_bytes': MAX_MESSAGE_BYTES, + 'server': 'ledmatrix-display', + } + return Response.success(request.id, dict(result), v=version) + + if request.v not in SUPPORTED_VERSIONS: + return Response.failure( + request.id, ErrorCode.UNSUPPORTED_VERSION, + f'protocol version {request.v} is not supported; ' + f'this display speaks {list(SUPPORTED_VERSIONS)}') + + try: + args = parse_args(request.cmd, request.args) + except ProtocolError as e: + return Response.failure(request.id, e.code, e.message, v=request.v) + + if request.cmd == Command.PING: + return Response.success(request.id, {'pong': True}, v=request.v) + + if request.cmd == Command.ON_DEMAND_STATUS: + if self._status_provider is None: + return Response.failure(request.id, ErrorCode.INTERNAL, 'no status available', + v=request.v) + return Response.success(request.id, self._status_provider(), v=request.v) + + if request.cmd in QUEUED_COMMANDS and isinstance(args, (OnDemandStartArgs, + OnDemandStopArgs)): + command = QueuedCommand(request_id=request.id, cmd=request.cmd, args=args, + received_at=time.time(), + peer_uid=peer.uid if peer is not None else None) + try: + self._queue.put_nowait(command) + except queue.Full: + logger.warning("Control socket queue full; refusing %s %s", + request.cmd, request.id) + return Response.failure(request.id, ErrorCode.BUSY, + 'the display is not taking commands right now', + v=request.v) + self._pending.set() + ack: AckResult = {'accepted': True, 'request_id': request.id, + 'queued': self._queue.qsize()} + logger.info("Control socket accepted %s %s", request.cmd, request.id) + return Response.success(request.id, dict(ack), v=request.v) + + # A command in COMMANDS with no handler here is a bug in this module. + return Response.failure(request.id, ErrorCode.INTERNAL, + f'{request.cmd} is not implemented', v=request.v) + + +def start_control_server(status_provider: Optional[StatusProvider] = None, + cache_dir: Optional[str] = None, + environ: Optional[Mapping[str, str]] = None) -> Optional[ControlServer]: + """Start the display's control socket, or return None when it can't run. + + None covers Windows, ``LEDMATRIX_CONTROL_SOCKET=off`` and any failure to + bind; in every case the web interface falls back to the file mailbox. + """ + path = server_socket_path(environ) + if path is None: + logger.debug("Control socket disabled or unsupported here; using the file mailbox only") + return None + server = ControlServer(path, status_provider, resolve_socket_group(cache_dir)) + return server if server.start() else None + + +__all__ = [ + 'ControlServer', 'PeerCredentials', 'QueuedCommand', 'StatusProvider', + 'peer_allowed', 'peer_credentials', 'process_groups', 'resolve_socket_group', + 'server_socket_path', 'start_control_server', 'PROTOCOL_VERSION', +] diff --git a/test/conftest.py b/test/conftest.py index 36901316..99d775f3 100644 --- a/test/conftest.py +++ b/test/conftest.py @@ -315,6 +315,19 @@ def _hermetic_display_watchdog(monkeypatch): display_watchdog.RenderWatchdog(environ={}, heartbeat_dir=None)) +@pytest.fixture(autouse=True) +def _hermetic_control_socket(monkeypatch): + """Keep the control socket (src/ipc) off the host. + + DisplayController.run() would serve /run/ledmatrix/control.sock -- or + find the live display's already there, when the suite runs on a device + -- and the web routes would send on-demand commands to that display. + Off by default; the socket tests point it at a tmp_path of their own. + """ + from src.ipc.contract import SOCKET_PATH_ENV + monkeypatch.setenv(SOCKET_PATH_ENV, 'off') + + @pytest.fixture(autouse=True) def reset_logging(): """Reset logging configuration before each test.""" diff --git a/test/test_api_v3_on_demand_socket.py b/test/test_api_v3_on_demand_socket.py new file mode 100644 index 00000000..8498b8bd --- /dev/null +++ b/test/test_api_v3_on_demand_socket.py @@ -0,0 +1,206 @@ +"""POST /display/on-demand/start and /stop: control socket first, mailbox fallback. + +The routes hand the request to the display over the control socket +(src/ipc) and get an acknowledgement. On any failure -- no socket (a stopped +display, or one older than the socket), a timeout, a refusal, a bug in the +client -- they write the file mailbox exactly as they did before the socket +existed. These tests pin both paths, that exactly one of them is used, that +the response says which, and that the request id is the same either way (the +display deduplicates on it). + +The socket client is patched at the route's module attribute; the last class +runs a real server on a temp socket (Linux/macOS only). +""" + +import os +import sys +from pathlib import Path +from unittest.mock import patch + +import pytest + +sys.path.insert(0, str(Path(__file__).parent.parent)) + +from test._api_v3_test_helpers import api_v3_client, api_v3_module # noqa: F401,E402 + +from src.ipc import client as control_client # noqa: E402 +from src.ipc import contract as c # noqa: E402 + +START_URL = "/api/v3/display/on-demand/start" +STOP_URL = "/api/v3/display/on-demand/stop" +MAILBOX = "display_on_demand_request" +CLIENT = "web_interface.blueprints.api_v3.display.control_client" + + +@pytest.fixture +def service(api_v3_module): + """A running display service; records systemctl calls and mailbox writes.""" + api_v3_module.api_v3.plugin_catalog = None + api_v3_module.api_v3.config_manager = None + state = {"active": True} + calls = [] + + def status(): + return {"active": state["active"]} + + def systemctl(args): + calls.append(("systemctl", args[-2])) + if args[-2:] == ["start", "ledmatrix.service"]: + state["active"] = True + return {"returncode": 0, "stdout": "", "stderr": ""} + + cache = api_v3_module.api_v3.cache_manager + cache.set.side_effect = lambda key, value, *a, **kw: calls.append(("cache", key)) + with patch("web_interface.blueprints.api_v3._get_display_service_status", + side_effect=status), \ + patch("web_interface.blueprints.api_v3.display._get_display_service_status", + side_effect=status), \ + patch("web_interface.blueprints.api_v3._run_systemctl_command", + side_effect=systemctl), \ + patch("web_interface.blueprints.api_v3.display._stop_display_service"): + yield {"state": state, "cache": cache, "calls": calls} + + +def _mailbox_writes(cache): + return [call.args[1] for call in cache.set.call_args_list + if call.args and call.args[0] == MAILBOX] + + +def _ack(request_id, *a, **kw): + return {"accepted": True, "request_id": request_id, "queued": 1} + + +class TestSocketPath: + def test_start_goes_over_the_socket_and_skips_the_mailbox(self, api_v3_client, service): + with patch(f"{CLIENT}.on_demand_start", side_effect=_ack) as start: + resp = api_v3_client.post(START_URL, json={ + "plugin_id": "weather", "mode": "weather_current", + "duration": 60, "pinned": True}) + assert resp.status_code == 200, resp.get_json() + data = resp.get_json()["data"] + assert data["transport"] == "socket" + assert "socket_error" not in data + assert _mailbox_writes(service["cache"]) == [] + start.assert_called_once_with(data["request_id"], "weather", "weather_current", 60, True) + + def test_a_callers_request_id_is_passed_through(self, api_v3_client, service): + with patch(f"{CLIENT}.on_demand_start", side_effect=_ack) as start: + data = api_v3_client.post(START_URL, json={ + "plugin_id": "weather", "request_id": "ha-123"}).get_json()["data"] + assert data["request_id"] == "ha-123" + assert start.call_args.args[0] == "ha-123" + + def test_stop_goes_over_the_socket(self, api_v3_client, service): + with patch(f"{CLIENT}.on_demand_stop", side_effect=_ack) as stop: + data = api_v3_client.post(STOP_URL, json={}).get_json()["data"] + assert data["transport"] == "socket" + stop.assert_called_once_with(data["request_id"]) + assert _mailbox_writes(service["cache"]) == [] + + +class TestMailboxFallback: + @pytest.mark.parametrize("reason", [ + "no_socket", "refused", "timeout", "closed", "bad_response", "invalid_request", + "busy", "unknown_command", "unsupported_version", "disabled", "unsupported", + ]) + def test_any_socket_failure_writes_the_mailbox_as_before( + self, api_v3_client, service, reason): + with patch(f"{CLIENT}.on_demand_start", + side_effect=control_client.ControlError(reason, "x")): + resp = api_v3_client.post(START_URL, json={ + "plugin_id": "weather", "mode": "weather_current", + "duration": 60, "pinned": True}) + assert resp.status_code == 200 + data = resp.get_json()["data"] + assert data["transport"] == "mailbox" + assert data["socket_error"] == reason + [write] = _mailbox_writes(service["cache"]) + assert write["request_id"] == data["request_id"] + assert write["action"] == "start" + assert (write["plugin_id"], write["mode"], write["duration"], write["pinned"]) == \ + ("weather", "weather_current", 60, True) + + def test_a_client_bug_still_falls_back(self, api_v3_client, service): + with patch(f"{CLIENT}.on_demand_start", side_effect=RuntimeError("boom")): + resp = api_v3_client.post(START_URL, json={"plugin_id": "weather"}) + assert resp.status_code == 200 + assert resp.get_json()["data"]["socket_error"] == "internal" + assert len(_mailbox_writes(service["cache"])) == 1 + + def test_an_unknown_reason_is_reported_as_other(self, api_v3_client, service): + # Only known codes are echoed back; anything else stays server-side. + with patch(f"{CLIENT}.on_demand_start", + side_effect=control_client.ControlError("/run/secret/path", "x")): + data = api_v3_client.post(START_URL, json={"plugin_id": "weather"}).get_json()["data"] + assert data["transport"] == "mailbox" + assert data["socket_error"] == "other" + assert len(_mailbox_writes(service["cache"])) == 1 + + def test_every_display_error_code_is_reportable(self): + from web_interface.blueprints.api_v3 import display + codes = {v for k, v in vars(c.ErrorCode).items() if not k.startswith("_")} + assert codes <= set(display._REPORTABLE_SOCKET_REASONS) + + def test_stop_falls_back(self, api_v3_client, service): + with patch(f"{CLIENT}.on_demand_stop", + side_effect=control_client.ControlError("timeout")): + data = api_v3_client.post(STOP_URL, json={}).get_json()["data"] + assert data["transport"] == "mailbox" + [write] = _mailbox_writes(service["cache"]) + assert write == {"request_id": data["request_id"], "action": "stop", + "timestamp": write["timestamp"]} + + def test_a_stopped_display_gets_the_mailbox_before_it_is_started( + self, api_v3_client, service): + service["state"]["active"] = False + with patch(f"{CLIENT}.on_demand_start", + side_effect=control_client.ControlError("no_socket")): + resp = api_v3_client.post(START_URL, json={"plugin_id": "weather"}) + assert resp.status_code == 200 + assert service["calls"] == [("cache", MAILBOX), ("systemctl", "start")] + + def test_the_socket_is_off_in_the_test_suite(self, api_v3_client, service): + # conftest's _hermetic_control_socket: a suite run on a device must + # not drive the live display. + assert os.environ[c.SOCKET_PATH_ENV] == "off" + data = api_v3_client.post(START_URL, json={"plugin_id": "weather"}).get_json()["data"] + assert data["transport"] == "mailbox" + assert data["socket_error"] in ("disabled", "unsupported") # Linux, Windows + + +@pytest.mark.skipif(not c.socket_supported(), reason="AF_UNIX sockets are Linux/macOS only") +class TestRealSocket: + @pytest.fixture + def live(self, monkeypatch): + import shutil + import tempfile + from src.ipc.server import ControlServer + d = tempfile.mkdtemp(prefix="lmipc-") + path = os.path.join(d, "control.sock") + server = ControlServer(path, status_provider=dict) + assert server.start() + monkeypatch.setenv(c.SOCKET_PATH_ENV, path) + yield server + server.close() + shutil.rmtree(d, ignore_errors=True) + + def test_start_is_acked_and_queued(self, api_v3_client, service, live): + data = api_v3_client.post(START_URL, json={ + "plugin_id": "weather", "duration": "30"}).get_json()["data"] + assert data["transport"] == "socket" + [cmd] = live.drain() + payload = cmd.as_on_demand_request() + assert payload["request_id"] == data["request_id"] + assert payload["plugin_id"] == "weather" and payload["duration"] == 30.0 + assert _mailbox_writes(service["cache"]) == [] + + def test_stop_is_acked_and_queued(self, api_v3_client, service, live): + data = api_v3_client.post(STOP_URL, json={}).get_json()["data"] + assert data["transport"] == "socket" + assert [x.request_id for x in live.drain()] == [data["request_id"]] + + def test_a_display_that_went_away_falls_back(self, api_v3_client, service, live): + live.close() + data = api_v3_client.post(START_URL, json={"plugin_id": "weather"}).get_json()["data"] + assert data["transport"] == "mailbox" and data["socket_error"] == "no_socket" + assert len(_mailbox_writes(service["cache"])) == 1 diff --git a/test/test_ipc_contract.py b/test/test_ipc_contract.py new file mode 100644 index 00000000..57be30b1 --- /dev/null +++ b/test/test_ipc_contract.py @@ -0,0 +1,275 @@ +"""The control socket's contract (src/ipc/contract.py): messages and framing. + +Pure data, so every test here runs on every platform. What they pin: + +* a request and a response survive encode -> decode -> parse unchanged, and + the on-demand arguments carry exactly what the file mailbox carries; +* the envelope and the arguments refuse what the display could not act on + (missing ids, wrong types, a non-finite duration) with a stable error code; +* framing never holds more than one message's worth of bytes, however the + bytes arrive; +* where the socket is looked for, and how it is switched off. +""" + +import json +import math + +import pytest + +from src.ipc import contract as c +from src.ipc.contract import ( + Command, ErrorCode, FrameReader, OnDemandStartArgs, OnDemandStopArgs, + ProtocolError, Request, Response, +) + + +def _wire(obj): + """Encode then decode, as one side's bytes reach the other.""" + data = c.encode_message(obj) + assert data.endswith(b'\n') and data.count(b'\n') == 1 + return c.decode_message(data) + + +class TestRoundTrip: + def test_request(self): + req = Request(id='abc-1', cmd=Command.ON_DEMAND_START, + args={'plugin_id': 'clock', 'mode': None, 'duration': 30.0, + 'pinned': True}) + back = Request.from_dict(_wire(req.to_dict())) + assert back == req + assert back.v == c.PROTOCOL_VERSION + + def test_success_response(self): + resp = Response.success('abc-1', {'accepted': True, 'request_id': 'abc-1', 'queued': 1}) + back = Response.from_dict(_wire(resp.to_dict())) + assert back == resp + assert back.ok and back.error is None + + def test_failure_response(self): + resp = Response.failure('abc-1', ErrorCode.BUSY, 'queue full') + wire = _wire(resp.to_dict()) + assert wire == {'v': 1, 'id': 'abc-1', 'ok': False, + 'error': {'code': 'busy', 'message': 'queue full'}} + assert Response.from_dict(wire) == resp + + def test_failure_without_an_id(self): + wire = _wire(Response.failure(None, ErrorCode.BAD_JSON, 'nope').to_dict()) + assert wire['id'] is None + assert Response.from_dict(wire).id is None + + def test_start_args_round_trip(self): + args = OnDemandStartArgs(plugin_id='clock', mode='clock_main', duration=45.0, + pinned=True) + assert OnDemandStartArgs.from_dict(_wire(args.to_dict())) == args + + def test_encoded_messages_are_ascii_single_lines(self): + data = c.encode_message({'v': 1, 'id': 'x', 'cmd': 'ping', + 'args': {'text': 'line1\nline2 café'}}) + assert data.count(b'\n') == 1 + data.decode('ascii') + assert c.decode_message(data)['args']['text'] == 'line1\nline2 café' + + +class TestEnvelopeValidation: + @pytest.mark.parametrize('obj', [ + [], 'x', 1, None, + ]) + def test_not_an_object(self, obj): + with pytest.raises(ProtocolError) as e: + Request.from_dict(obj) + assert e.value.code == ErrorCode.BAD_REQUEST + + @pytest.mark.parametrize('bad_id', [None, '', 7, 'x' * (c.MAX_ID_LENGTH + 1), 'a\nb']) + def test_bad_id(self, bad_id): + with pytest.raises(ProtocolError) as e: + Request.from_dict({'v': 1, 'id': bad_id, 'cmd': 'ping'}) + assert e.value.code == ErrorCode.BAD_REQUEST + assert e.value.request_id is None + + @pytest.mark.parametrize('v', [None, '1', 1.0, True]) + def test_bad_version_type_keeps_the_id(self, v): + with pytest.raises(ProtocolError) as e: + Request.from_dict({'v': v, 'id': 'r1', 'cmd': 'ping'}) + assert e.value.code == ErrorCode.BAD_REQUEST + assert e.value.request_id == 'r1' + + def test_missing_cmd(self): + with pytest.raises(ProtocolError) as e: + Request.from_dict({'v': 1, 'id': 'r1'}) + assert e.value.code == ErrorCode.BAD_REQUEST + + def test_args_default_to_empty(self): + assert Request.from_dict({'v': 1, 'id': 'r', 'cmd': 'ping'}).args == {} + assert Request.from_dict({'v': 1, 'id': 'r', 'cmd': 'ping', 'args': None}).args == {} + + def test_args_must_be_an_object(self): + with pytest.raises(ProtocolError) as e: + Request.from_dict({'v': 1, 'id': 'r', 'cmd': 'ping', 'args': [1]}) + assert e.value.code == ErrorCode.BAD_REQUEST + + def test_an_unknown_version_parses(self): + # The server, not the parser, decides about versions, so that hello + # can negotiate. + assert Request.from_dict({'v': 99, 'id': 'r', 'cmd': 'hello'}).v == 99 + + @pytest.mark.parametrize('obj', [ + {'v': 1, 'id': 'r', 'ok': 'yes'}, + {'v': 1, 'id': 'r', 'ok': False}, + {'v': 1, 'id': 'r', 'ok': False, 'error': {'message': 'x'}}, + {'v': 1, 'id': 5, 'ok': True}, + {'v': 1, 'id': 'r', 'ok': True, 'result': [1]}, + {'id': 'r', 'ok': True}, + ]) + def test_malformed_responses(self, obj): + with pytest.raises(ProtocolError): + Response.from_dict(obj) + + +class TestOnDemandArgs: + def test_plugin_or_mode_is_required(self): + with pytest.raises(ProtocolError) as e: + OnDemandStartArgs.from_dict({'duration': 10}) + assert e.value.code == ErrorCode.INVALID_ARGS + + def test_mode_alone_is_enough(self): + assert OnDemandStartArgs.from_dict({'mode': 'nfl_live'}).mode == 'nfl_live' + + @pytest.mark.parametrize('raw, seconds', [ + (None, None), ('', None), (0, None), (45, 45.0), (2.5, 2.5), ('30', 30.0), + ]) + def test_duration(self, raw, seconds): + assert OnDemandStartArgs.from_dict({'plugin_id': 'p', 'duration': raw}).duration == seconds + + @pytest.mark.parametrize('raw', [-1, 'soon', True, [5], math.inf, math.nan, 'inf']) + def test_bad_duration(self, raw): + with pytest.raises(ProtocolError) as e: + OnDemandStartArgs.from_dict({'plugin_id': 'p', 'duration': raw}) + assert e.value.code == ErrorCode.INVALID_ARGS + + @pytest.mark.parametrize('pinned', ['true', 1, 'false']) + def test_pinned_must_be_a_real_boolean(self, pinned): + # The web route coerces "false" to False before it gets here; the + # contract does not guess (bool("false") is True). + with pytest.raises(ProtocolError): + OnDemandStartArgs.from_dict({'plugin_id': 'p', 'pinned': pinned}) + + @pytest.mark.parametrize('name', [5, 'x' * (c.MAX_NAME_LENGTH + 1), 'a\nb']) + def test_bad_names(self, name): + with pytest.raises(ProtocolError): + OnDemandStartArgs.from_dict({'plugin_id': name}) + + def test_unknown_command(self): + with pytest.raises(ProtocolError) as e: + c.parse_args('reboot', {}) + assert e.value.code == ErrorCode.UNKNOWN_COMMAND + + @pytest.mark.parametrize('cmd', c.COMMANDS) + def test_every_command_has_an_argument_type(self, cmd): + args = {'plugin_id': 'p'} if cmd == Command.ON_DEMAND_START else {} + c.parse_args(cmd, args) + + def test_hello_versions(self): + assert c.HelloArgs.from_dict({'versions': [1, 2], 'client': 'web'}).versions == (1, 2) + for bad in ([], ['1'], 'x', [True]): + with pytest.raises(ProtocolError): + c.HelloArgs.from_dict({'versions': bad}) + + def test_negotiation(self): + assert c.negotiate_version((1,)) == 1 + assert c.negotiate_version((1, 7)) == 1 + assert c.negotiate_version((7,)) is None + + +class TestMailboxShape: + """Socket commands are handed to the mailbox's own handler, so they must + look exactly like what the web route writes to the mailbox.""" + + def test_start(self): + args = OnDemandStartArgs(plugin_id='clock', mode='clock_main', duration=60.0, + pinned=True) + payload = c.on_demand_request('rid', args, 123.0) + assert payload == {'request_id': 'rid', 'action': 'start', 'plugin_id': 'clock', + 'mode': 'clock_main', 'duration': 60.0, 'pinned': True, + 'timestamp': 123.0, 'source': 'socket'} + + def test_stop(self): + payload = c.on_demand_request('rid', OnDemandStopArgs(), 5.0) + assert payload['action'] == 'stop' and payload['request_id'] == 'rid' + + +class TestFraming: + def test_one_message_in_pieces(self): + data = c.encode_message({'v': 1, 'id': 'a', 'cmd': 'ping'}) + reader = FrameReader() + out = [] + for i in range(len(data)): + out += reader.feed(data[i:i + 1]) + assert [json.loads(x) for x in out] == [{'v': 1, 'id': 'a', 'cmd': 'ping'}] + assert reader.pending == 0 + + def test_several_messages_in_one_chunk(self): + data = b''.join(c.encode_message({'n': n}) for n in range(3)) + assert [json.loads(x)['n'] for x in FrameReader().feed(data)] == [0, 1, 2] + + def test_blank_lines_are_skipped(self): + assert FrameReader().feed(b'\n\r\n \n{"a":1}\n') == [b'{"a":1}'] + + def test_a_line_over_the_limit_is_refused(self): + reader = FrameReader(max_bytes=32) + with pytest.raises(ProtocolError) as e: + reader.feed(b'x' * 40 + b'\n') + assert e.value.code == ErrorCode.MESSAGE_TOO_LARGE + + def test_a_line_that_never_ends_is_refused_at_the_limit(self): + reader = FrameReader(max_bytes=32) + reader.feed(b'x' * 31) + with pytest.raises(ProtocolError) as e: + reader.feed(b'x') + assert e.value.code == ErrorCode.MESSAGE_TOO_LARGE + + def test_exactly_the_limit_is_allowed(self): + reader = FrameReader(max_bytes=8) + assert reader.feed(b'1234567\n') == [b'1234567'] + + def test_encode_refuses_an_oversize_message(self): + with pytest.raises(ProtocolError) as e: + c.encode_message({'blob': 'x' * c.MAX_MESSAGE_BYTES}) + assert e.value.code == ErrorCode.MESSAGE_TOO_LARGE + + def test_encode_refuses_non_json(self): + with pytest.raises(ProtocolError): + c.encode_message({'x': math.nan}) + with pytest.raises(ProtocolError): + c.encode_message({'x': object()}) + + @pytest.mark.parametrize('line', [b'{', b'[1,2]', b'"x"', b'\xff\xfe', b'null']) + def test_decode_garbage(self, line): + with pytest.raises(ProtocolError) as e: + c.decode_message(line) + assert e.value.code == ErrorCode.BAD_JSON + + +class TestSocketLocation: + def test_default(self): + paths = c.client_socket_paths({}) + assert paths[0] == '/run/ledmatrix/control.sock' + assert len(paths) == 2 and paths[1].endswith('control.sock') + + def test_configured_path_is_the_only_one_tried(self): + assert c.client_socket_paths({c.SOCKET_PATH_ENV: '/x/y.sock'}) == ['/x/y.sock'] + + @pytest.mark.parametrize('value', ['off', 'OFF', '0', 'false', 'disabled', ' none ']) + def test_switched_off(self, value): + env = {c.SOCKET_PATH_ENV: value} + assert c.socket_disabled(env) + assert c.client_socket_paths(env) == [] + assert c.configured_socket_path(env) is None + + def test_dev_path_is_per_user(self): + assert c.dev_socket_path(1000) != c.dev_socket_path(1001) + assert 'ledmatrix-1000' in c.dev_socket_path(1000) + + def test_the_default_dir_is_the_heartbeats(self): + # One RuntimeDirectory= serves both (#687). + from src import display_watchdog + assert c.DEFAULT_SOCKET_DIR == display_watchdog.HEARTBEAT_DIR diff --git a/test/test_ipc_display_hook.py b/test/test_ipc_display_hook.py new file mode 100644 index 00000000..1afffe50 --- /dev/null +++ b/test/test_ipc_display_hook.py @@ -0,0 +1,229 @@ +"""DisplayController's side of the control socket. + +The server's handlers only queue; the render thread drains the queue where +it reads the file mailbox (_poll_on_demand_requests) and hands each command +to the mailbox's own handler (_handle_on_demand_request). These tests pin +that hook: + +* a socket command is applied by the same code as a mailbox request, with + its request id, and without waiting for the mailbox's 0.25 s read floor; +* a request that arrives both ways (a client that timed out after the + command was queued, then wrote the mailbox) is activated once; +* a command that fails is contained, and the ones after it still run; +* cleanup closes the socket; a disabled socket changes nothing. +""" + +import os +import time +from unittest.mock import MagicMock + +import pytest + +from src.ipc import client +from src.ipc import contract as c +from src.ipc.contract import Command, OnDemandStartArgs, OnDemandStopArgs +from src.ipc.server import QueuedCommand + + +def _start(rid, plugin_id='clock', **kw): + return QueuedCommand(request_id=rid, cmd=Command.ON_DEMAND_START, + args=OnDemandStartArgs(plugin_id=plugin_id, **kw), + received_at=time.time()) + + +def _stop(rid): + return QueuedCommand(request_id=rid, cmd=Command.ON_DEMAND_STOP, + args=OnDemandStopArgs(), received_at=time.time()) + + +class FakeServer: + def __init__(self, *commands): + self.commands = list(commands) + self.closed = False + + @property + def has_pending(self): + return bool(self.commands) + + def drain(self): + out, self.commands = self.commands, [] + return out + + def close(self): + self.closed = True + + +@pytest.fixture +def controller(test_display_controller): + c_ = test_display_controller + c_.on_demand_active = False + c_.on_demand_request_id = None + c_._last_on_demand_poll = None + mailbox = {'value': None} + + def fake_get(key, *a, **kw): + if key == 'display_on_demand_request': + return mailbox['value'] + return None + + def fake_delete(key): + if key == 'display_on_demand_request': + mailbox['value'] = None + + c_.cache_manager.get = MagicMock(side_effect=fake_get) + c_.cache_manager.set = MagicMock() + c_.cache_manager.delete = MagicMock(side_effect=fake_delete) + c_._activate_on_demand = MagicMock() + c_.mailbox = mailbox + return c_ + + +class TestDrain: + def test_a_socket_start_goes_through_the_mailbox_handler(self, controller): + controller._control_server = FakeServer(_start('sock-1', duration=30.0, pinned=True)) + controller._poll_on_demand_requests() + controller._activate_on_demand.assert_called_once() + request = controller._activate_on_demand.call_args.args[0] + assert request['request_id'] == 'sock-1' + assert request['action'] == 'start' + assert request['plugin_id'] == 'clock' + assert request['duration'] == 30.0 and request['pinned'] is True + assert controller.on_demand_request_id == 'sock-1' + # The same restart-replay guard as a mailbox request. + controller.cache_manager.set.assert_any_call( + 'display_on_demand_processed_id', 'sock-1', ttl=3600) + + def test_socket_commands_skip_the_mailbox_floor(self, controller): + server = FakeServer() + controller._control_server = server + controller._poll_on_demand_requests() # reads the mailbox, sets the floor + reads = controller.cache_manager.get.call_count + server.commands.append(_start('quick')) + controller._poll_on_demand_requests() # within the floor + controller._activate_on_demand.assert_called_once() + mailbox_reads = [call for call in controller.cache_manager.get.call_args_list[reads:] + if call.args[0] == 'display_on_demand_request'] + # Only _consume_on_demand_request's compare-before-delete re-read. + assert len(mailbox_reads) <= 1 + + def test_a_request_that_came_both_ways_is_activated_once(self, controller): + controller._control_server = FakeServer(_start('dup')) + controller.mailbox['value'] = {'request_id': 'dup', 'action': 'start', + 'plugin_id': 'clock'} + controller._poll_on_demand_requests() + controller._last_on_demand_poll = None + controller._poll_on_demand_requests() + controller._activate_on_demand.assert_called_once() + assert controller.mailbox['value'] is None, "the duplicate was left in the mailbox" + + def test_a_fallback_write_landing_later_is_ignored(self, controller): + controller._control_server = FakeServer(_start('late')) + controller._poll_on_demand_requests() + controller.mailbox['value'] = {'request_id': 'late', 'action': 'start', + 'plugin_id': 'clock'} + controller._last_on_demand_poll = None + controller._poll_on_demand_requests() + controller._activate_on_demand.assert_called_once() + + def test_the_mailbox_still_works_alongside(self, controller): + controller._control_server = FakeServer() + controller.mailbox['value'] = {'request_id': 'mb', 'action': 'start', 'plugin_id': 'p'} + controller._poll_on_demand_requests() + assert controller._activate_on_demand.call_args.args[0]['request_id'] == 'mb' + + def test_a_socket_stop_ends_on_demand(self, controller): + controller.on_demand_active = True + controller._clear_on_demand = MagicMock() + controller._control_server = FakeServer(_stop('halt')) + controller._poll_on_demand_requests() + controller._clear_on_demand.assert_called_once_with(reason='requested-stop') + + def test_commands_run_in_arrival_order(self, controller): + seen = [] + controller._activate_on_demand = MagicMock( + side_effect=lambda r: seen.append(r['request_id'])) + controller._control_server = FakeServer(_start('a'), _start('b'), _start('c')) + controller._poll_on_demand_requests() + assert seen == ['a', 'b', 'c'] + + def test_a_failing_command_is_contained(self, controller): + calls = [] + + def activate(request): + calls.append(request['request_id']) + if request['request_id'] == 'bad': + raise RuntimeError('plugin exploded') + + controller._activate_on_demand = MagicMock(side_effect=activate) + controller._control_server = FakeServer(_start('bad'), _start('good')) + controller._poll_on_demand_requests() + assert calls == ['bad', 'good'] + + def test_no_server_means_mailbox_only(self, controller): + controller._control_server = None + controller._poll_on_demand_requests() + controller._activate_on_demand.assert_not_called() + + +class TestPendingChangesFloor: + def test_a_queued_command_skips_the_floor(self, controller): + server = FakeServer() + controller._control_server = server + controller._service_pending_changes() + server.commands.append(_start('now')) + controller._service_pending_changes() # well inside the 0.25 s floor + controller._activate_on_demand.assert_called_once() + + def test_nothing_queued_keeps_the_floor(self, controller): + controller._control_server = FakeServer() + controller._poll_on_demand_requests = MagicMock() + controller._service_pending_changes() + controller._service_pending_changes() + assert controller._poll_on_demand_requests.call_count == 1 + + +class TestLifecycle: + def test_status_snapshot(self, controller): + controller.current_display_mode = 'clock_main' + controller.on_demand_active = True + controller.on_demand_plugin_id = 'clock' + controller.on_demand_expires_at = None + status = controller._control_status() + assert status['current_mode'] == 'clock_main' + assert status['on_demand']['active'] is True + assert status['on_demand']['plugin_id'] == 'clock' + c.encode_message(status) # it has to fit on the wire + + def test_cleanup_closes_the_socket(self, controller): + server = FakeServer() + controller._control_server = server + controller.cleanup() + assert server.closed + assert controller._control_server is None + + def test_disabled_socket_starts_nothing(self, controller): + # conftest sets LEDMATRIX_CONTROL_SOCKET=off for every test. + controller._start_control_server() + assert controller._control_server is None + + @pytest.mark.skipif(not c.socket_supported(), reason='AF_UNIX sockets are Linux/macOS only') + def test_end_to_end(self, controller, monkeypatch): + import shutil + import tempfile + d = tempfile.mkdtemp(prefix='lmipc-') + path = os.path.join(d, 'control.sock') + monkeypatch.setenv(c.SOCKET_PATH_ENV, path) + try: + controller._start_control_server() + assert controller._control_server is not None + ack = client.on_demand_start('e2e', 'clock', None, 15, False, paths=[path]) + assert ack['accepted'] is True and ack['request_id'] == 'e2e' + status = client.on_demand_status(paths=[path]) + assert 'on_demand' in status and 'current_mode' in status + controller._service_pending_changes() + request = controller._activate_on_demand.call_args.args[0] + assert request['request_id'] == 'e2e' and request['duration'] == 15.0 + controller.cleanup() + assert not os.path.exists(path) + finally: + shutil.rmtree(d, ignore_errors=True) diff --git a/test/test_ipc_server.py b/test/test_ipc_server.py new file mode 100644 index 00000000..4d15e47f --- /dev/null +++ b/test/test_ipc_server.py @@ -0,0 +1,570 @@ +"""The display side of the control socket (src/ipc/server.py). + +Two layers: + +* ``handle_line`` and the permission model are plain functions of their + input, tested on every platform: every request gets exactly one answer, + garbage is answered rather than raised, queued commands are acked with + their request id, and the render thread drains them in order. +* The socket itself (``TestLiveSocket``, ``TestPermissions``) needs AF_UNIX, + so those tests are skipped on Windows and run on Linux (CI, WSL, a Pi): a + real server on a tmp_path socket, driven by the real client and by raw + sockets that misbehave -- garbage, oversize lines, a client that hangs up + mid-message, one that never finishes -- while the server keeps serving. +""" + +import json +import os +import socket +import stat +import threading +import time + +import pytest + +from src.ipc import client +from src.ipc import contract as c +from src.ipc import server as srv +from src.ipc.contract import Command, ErrorCode +from src.ipc.server import ControlServer, PeerCredentials, peer_allowed + +needs_unix_sockets = pytest.mark.skipif(not c.socket_supported(), + reason='AF_UNIX sockets are Linux/macOS only') + + +def _line(obj): + return json.dumps(obj).encode() + + +def _req(cmd, args=None, rid='r1', v=1): + return _line({'v': v, 'id': rid, 'cmd': cmd, 'args': args or {}}) + + +@pytest.fixture +def status(): + return {'on_demand': {'active': False, 'status': 'idle'}, 'current_mode': 'clock'} + + +@pytest.fixture +def server(status, tmp_path): + """A server that is never started: handle_line and drain only.""" + return ControlServer(str(tmp_path / 'unused.sock'), status_provider=lambda: dict(status), + queue_size=3) + + +class TestHandleLine: + def test_ping(self, server): + resp = server.handle_line(_req(Command.PING)) + assert resp.ok and resp.id == 'r1' and resp.result == {'pong': True} + + def test_hello_negotiates(self, server): + resp = server.handle_line(_req(Command.HELLO, {'versions': [1, 5], 'client': 't'}, v=5)) + assert resp.ok + assert resp.result['version'] == 1 + assert resp.result['commands'] == list(c.COMMANDS) + assert resp.result['max_message_bytes'] == c.MAX_MESSAGE_BYTES + assert resp.v == 1 + + def test_hello_with_nothing_in_common(self, server): + resp = server.handle_line(_req(Command.HELLO, {'versions': [9]}, v=9)) + assert not resp.ok and resp.error.code == ErrorCode.UNSUPPORTED_VERSION + + def test_other_commands_need_a_supported_version(self, server): + resp = server.handle_line(_req(Command.PING, v=2)) + assert not resp.ok and resp.error.code == ErrorCode.UNSUPPORTED_VERSION + assert resp.id == 'r1' + + @pytest.mark.parametrize('line, code, rid', [ + (b'not json', ErrorCode.BAD_JSON, None), + (b'[1]', ErrorCode.BAD_JSON, None), + (b'\xff', ErrorCode.BAD_JSON, None), + (_line({'v': 1, 'cmd': 'ping'}), ErrorCode.BAD_REQUEST, None), + (_line({'v': 'one', 'id': 'q', 'cmd': 'ping'}), ErrorCode.BAD_REQUEST, 'q'), + (_req('shutdown_the_pi'), ErrorCode.UNKNOWN_COMMAND, 'r1'), + (_req(Command.ON_DEMAND_START, {}), ErrorCode.INVALID_ARGS, 'r1'), + (_req(Command.ON_DEMAND_START, {'plugin_id': 'p', 'duration': 'x'}), + ErrorCode.INVALID_ARGS, 'r1'), + ]) + def test_garbage_is_answered_not_raised(self, server, line, code, rid): + resp = server.handle_line(line) + assert not resp.ok + assert resp.error.code == code + assert resp.id == rid + assert server.drain() == [] + + def test_start_is_queued_and_acked(self, server): + resp = server.handle_line(_req(Command.ON_DEMAND_START, + {'plugin_id': 'clock', 'duration': 30, 'pinned': True}, + rid='abc')) + assert resp.ok + assert resp.result == {'accepted': True, 'request_id': 'abc', 'queued': 1} + assert server.has_pending + [cmd] = server.drain() + assert not server.has_pending + payload = cmd.as_on_demand_request() + assert payload['request_id'] == 'abc' + assert payload['action'] == 'start' + assert payload['plugin_id'] == 'clock' + assert payload['duration'] == 30.0 and payload['pinned'] is True + + def test_stop_is_queued_and_acked(self, server): + resp = server.handle_line(_req(Command.ON_DEMAND_STOP, rid='s1')) + assert resp.ok and resp.result['request_id'] == 's1' + assert [x.as_on_demand_request()['action'] for x in server.drain()] == ['stop'] + + def test_drain_keeps_arrival_order(self, server): + for rid in ('a', 'b', 'c'): + server.handle_line(_req(Command.ON_DEMAND_START, {'plugin_id': 'p'}, rid=rid)) + assert [x.request_id for x in server.drain()] == ['a', 'b', 'c'] + assert server.drain() == [] + + def test_a_full_queue_says_busy_and_queues_nothing_more(self, server): + for rid in ('a', 'b', 'c'): + assert server.handle_line(_req(Command.ON_DEMAND_STOP, rid=rid)).ok + resp = server.handle_line(_req(Command.ON_DEMAND_STOP, rid='d')) + assert not resp.ok and resp.error.code == ErrorCode.BUSY and resp.id == 'd' + assert [x.request_id for x in server.drain()] == ['a', 'b', 'c'] + + def test_status_answers_from_the_provider_without_queueing(self, server, status): + status['on_demand']['active'] = True + resp = server.handle_line(_req(Command.ON_DEMAND_STATUS)) + assert resp.ok and resp.result['on_demand']['active'] is True + assert not server.has_pending + + def test_a_failing_status_provider_is_an_internal_error(self, tmp_path): + def boom(): + raise RuntimeError('render thread mid-update') + s = ControlServer(str(tmp_path / 'x.sock'), status_provider=boom) + resp = s.handle_line(_req(Command.ON_DEMAND_STATUS, rid='z')) + assert not resp.ok and resp.error.code == ErrorCode.INTERNAL and resp.id == 'z' + + +class TestPermissionModel: + """root, the display's own user, or the shared group -- nobody else.""" + OWN, GROUP = 0, 990 + + @pytest.mark.parametrize('cred, groups, allowed', [ + (PeerCredentials(1, 0, 0), None, True), # root + (PeerCredentials(1, 1000, 1000), frozenset({990}), True), # web user, in group + (PeerCredentials(1, 1000, 990), frozenset(), True), # primary group + (PeerCredentials(1, 1001, 1001), frozenset({27, 44}), False), + (PeerCredentials(1, 65534, 65534), frozenset(), False), # nobody + ]) + def test_model(self, cred, groups, allowed): + assert peer_allowed(cred, self.OWN, self.GROUP, groups) is allowed + + def test_own_user_without_a_group(self): + assert peer_allowed(PeerCredentials(1, 1000, 1000), 1000, None, frozenset()) + assert not peer_allowed(PeerCredentials(1, 1001, 1001), 1000, None, frozenset({1})) + + def test_group_database_decides_when_proc_is_unreadable(self): + seen = [] + + def in_group(uid, gid): + seen.append((uid, gid)) + return uid == 1000 + + cred = PeerCredentials(1, 1000, 1000) + assert peer_allowed(cred, 0, 990, None, in_group=in_group) + assert not peer_allowed(PeerCredentials(1, 1001, 1001), 0, 990, None, in_group=in_group) + assert seen == [(1000, 990), (1001, 990)] + + @pytest.mark.skipif(os.name != 'posix', reason='POSIX permission bits') + def test_socket_group_follows_a_shared_cache_dir(self, tmp_path): + shared = tmp_path / 'cache' + shared.mkdir() + os.chmod(shared, 0o2775) + assert srv.resolve_socket_group(str(shared)) == shared.stat().st_gid + + @pytest.mark.skipif(os.name != 'posix', reason='POSIX permission bits') + def test_a_private_cache_dir_falls_back_to_the_project_group(self, tmp_path, monkeypatch): + private = tmp_path / 'cache' + private.mkdir() + os.chmod(private, 0o755) + from src.common import permission_utils + monkeypatch.setattr(permission_utils, 'get_shared_group_gid', lambda: 4242) + assert srv.resolve_socket_group(str(private)) == 4242 + + def test_mode_is_group_only_with_a_group(self, tmp_path): + assert ControlServer(str(tmp_path / 'a'), group=990).socket_mode == 0o660 + assert ControlServer(str(tmp_path / 'b'), group=None).socket_mode == 0o600 + + +class TestWhereTheServerListens: + def test_off_means_no_server(self): + assert srv.server_socket_path({c.SOCKET_PATH_ENV: 'off'}) is None + assert srv.start_control_server(environ={c.SOCKET_PATH_ENV: 'off'}) is None + + @needs_unix_sockets + def test_configured(self): + assert srv.server_socket_path({c.SOCKET_PATH_ENV: '/tmp/x.sock'}) == '/tmp/x.sock' + + @needs_unix_sockets + def test_unprivileged_dev_run_uses_the_per_user_path(self, monkeypatch): + monkeypatch.setattr(os, 'geteuid', lambda: 1000) + monkeypatch.setattr(os, 'access', lambda p, m: False) + assert srv.server_socket_path({}) == c.dev_socket_path() + + @needs_unix_sockets + def test_root_uses_run(self, monkeypatch): + monkeypatch.setattr(os, 'geteuid', lambda: 0) + assert srv.server_socket_path({}) == c.DEFAULT_SOCKET_PATH + + @pytest.mark.skipif(c.socket_supported(), reason='Windows only') + def test_windows_skips_cleanly(self): + assert srv.server_socket_path({}) is None + assert ControlServer('x.sock').start() is False + with pytest.raises(client.ControlError) as e: + client.ping(paths=['x.sock']) + assert e.value.reason == 'unsupported' + + +# -- the real socket -------------------------------------------------------------------- + +@pytest.fixture +def sock_path(tmp_path_factory): + # AF_UNIX paths are limited to ~107 bytes; pytest's tmp_path can be longer. + import tempfile + d = tempfile.mkdtemp(prefix='lmipc-') + yield os.path.join(d, 'control.sock') + import shutil + shutil.rmtree(d, ignore_errors=True) + + +@pytest.fixture +def live(sock_path, status): + servers = [] + + def make(**kwargs): + kwargs.setdefault('status_provider', lambda: dict(status)) + s = ControlServer(sock_path, **kwargs) + assert s.start() + servers.append(s) + return s + + yield make + for s in servers: + s.close() + + +def _raw(path, timeout=2.0): + s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + s.settimeout(timeout) + s.connect(path) + return s + + +def _read_line(s): + buf = b'' + while not buf.endswith(b'\n'): + try: + chunk = s.recv(1) # one byte at a time: never eat the next line + except ConnectionResetError: + break + if not chunk: + break + buf += chunk + return json.loads(buf) if buf.endswith(b'\n') else None + + +@needs_unix_sockets +class TestLiveSocket: + def test_client_round_trip(self, live, sock_path): + live() + assert client.ping(paths=[sock_path]) == {'pong': True} + assert client.hello(paths=[sock_path])['version'] == 1 + assert client.on_demand_status(paths=[sock_path])['current_mode'] == 'clock' + + def test_ack_path(self, live, sock_path): + server = live() + ack = client.on_demand_start('req-1', 'clock', None, 20, True, paths=[sock_path]) + assert ack == {'accepted': True, 'request_id': 'req-1', 'queued': 1} + ack = client.on_demand_stop('req-2', paths=[sock_path]) + assert ack['request_id'] == 'req-2' + assert [(x.request_id, x.cmd) for x in server.drain()] == [ + ('req-1', Command.ON_DEMAND_START), ('req-2', Command.ON_DEMAND_STOP)] + + def test_socket_file_mode_and_cleanup(self, live, sock_path): + server = live(group=os.getgid()) + st = os.lstat(sock_path) + assert stat.S_ISSOCK(st.st_mode) + assert stat.S_IMODE(st.st_mode) == 0o660 + assert st.st_gid == os.getgid() + assert not [f for f in os.listdir(os.path.dirname(sock_path)) if f.endswith('.tmp')] + server.close() + assert not os.path.exists(sock_path) + + def test_without_a_group_only_the_owner_may_connect(self, live, sock_path): + live(group=None) + assert stat.S_IMODE(os.lstat(sock_path).st_mode) == 0o600 + + def test_garbage_then_a_good_request_on_one_connection(self, live, sock_path): + live() + s = _raw(sock_path) + try: + s.sendall(b'this is not json\n') + assert _read_line(s)['error']['code'] == ErrorCode.BAD_JSON + s.sendall(_req(Command.PING, rid='after') + b'\n') + resp = _read_line(s) + assert resp['ok'] and resp['id'] == 'after' + finally: + s.close() + + def test_two_requests_in_one_write(self, live, sock_path): + live() + s = _raw(sock_path) + try: + s.sendall(_req(Command.PING, rid='a') + b'\n' + _req(Command.PING, rid='b') + b'\n') + assert _read_line(s)['id'] == 'a' + assert _read_line(s)['id'] == 'b' + finally: + s.close() + + def test_oversize_is_refused_and_the_server_lives_on(self, live, sock_path): + live() + s = _raw(sock_path) + try: + try: + # Exactly the limit with no newline: the server has read it + # all when it refuses, so its answer is not lost to a reset. + s.sendall(b'{"pad":"' + b'x' * (c.MAX_MESSAGE_BYTES - 8)) + except OSError: + pass # the server may hang up before we finish writing + resp = _read_line(s) + assert resp['error']['code'] == ErrorCode.MESSAGE_TOO_LARGE + assert s.recv(10) == b'' # and hung up + finally: + s.close() + assert client.ping(paths=[sock_path]) == {'pong': True} + + def test_a_client_that_hangs_up_mid_message(self, live, sock_path): + server = live() + s = _raw(sock_path) + s.sendall(b'{"v":1,"id":"half","cmd":"on_demand.st') + s.close() + time.sleep(0.2) + assert client.ping(paths=[sock_path]) == {'pong': True} + assert server.drain() == [] + + def test_a_slow_client_is_dropped_and_blocks_nobody(self, live, sock_path): + live(io_timeout=0.2, message_timeout=0.5) + slow = _raw(sock_path, timeout=3) + try: + slow.sendall(b'{"v":1,') # ...and never finishes + t0 = time.monotonic() + assert client.ping(paths=[sock_path]) == {'pong': True} + assert time.monotonic() - t0 < 0.5, 'a slow client held up another' + assert slow.recv(100) == b'' # hung up on, not answered + finally: + slow.close() + + def test_an_idle_connection_is_closed(self, live, sock_path): + live(io_timeout=0.1, idle_timeout=0.3) + s = _raw(sock_path, timeout=3) + try: + assert s.recv(100) == b'' + finally: + s.close() + + def test_too_many_clients_are_told_busy(self, live, sock_path): + live(max_clients=2, io_timeout=0.2, idle_timeout=5) + held = [_raw(sock_path) for _ in range(2)] + try: + time.sleep(0.1) + extra = _raw(sock_path) + try: + assert _read_line(extra)['error']['code'] == ErrorCode.BUSY + finally: + extra.close() + finally: + for s in held: + s.close() + time.sleep(0.3) + assert client.ping(paths=[sock_path]) == {'pong': True} + + def test_many_concurrent_clients(self, live, sock_path): + server = live(queue_size=64) + errors = [] + + def go(n): + try: + client.on_demand_start(f'r{n}', 'p', None, paths=[sock_path], timeout=3) + except client.ControlError as e: # busy is allowed under load + if e.reason != ErrorCode.BUSY: + errors.append(e) + + threads = [threading.Thread(target=go, args=(n,)) for n in range(20)] + for t in threads: + t.start() + for t in threads: + t.join() + assert errors == [] + assert 0 < len(server.drain()) <= 20 + + def test_a_stale_socket_is_replaced(self, sock_path, status): + dead = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + dead.bind(sock_path) + dead.close() # file left behind, nothing listening + s = ControlServer(sock_path, status_provider=lambda: status) + try: + assert s.start() + assert client.ping(paths=[sock_path]) == {'pong': True} + finally: + s.close() + + def test_a_live_socket_is_not_stolen(self, live, sock_path, status): + live() + second = ControlServer(sock_path, status_provider=lambda: status) + assert second.start() is False + assert client.ping(paths=[sock_path]) == {'pong': True} + + def test_a_regular_file_is_never_removed(self, sock_path): + with open(sock_path, 'w') as f: + f.write('precious') + assert ControlServer(sock_path).start() is False + with open(sock_path) as f: + assert f.read() == 'precious' + + def test_close_leaves_a_successor_s_socket_alone(self, sock_path, status): + first = ControlServer(sock_path, status_provider=lambda: status) + assert first.start() + first._close_socket() # dead, but still owns the path + os.unlink(sock_path) + second = ControlServer(sock_path, status_provider=lambda: status) + assert second.start() + try: + first.close() # must not unlink second's file + assert client.ping(paths=[sock_path]) == {'pong': True} + finally: + second.close() + + def test_client_reasons(self, sock_path, tmp_path): + with pytest.raises(client.ControlError) as e: + client.ping(paths=[sock_path]) + assert e.value.reason == 'no_socket' + dead = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + dead.bind(sock_path) + try: + with pytest.raises(client.ControlError) as e: + client.ping(paths=[sock_path]) + assert e.value.reason == 'refused' + finally: + dead.close() + with pytest.raises(client.ControlError) as e: + client.ping(paths=[]) + assert e.value.reason == 'disabled' + with pytest.raises(client.ControlError) as e: + client.on_demand_start('x', None, None, paths=[sock_path]) + assert e.value.reason == 'invalid_request' + + def test_a_display_that_never_answers_times_out(self, sock_path): + mute = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + mute.bind(sock_path) + mute.listen(1) # accepts at the kernel, never replies + try: + t0 = time.monotonic() + with pytest.raises(client.ControlError) as e: + client.ping(paths=[sock_path], timeout=0.3) + assert e.value.reason == 'timeout' + assert time.monotonic() - t0 < 1.0 + finally: + mute.close() + + def test_the_dev_directory_must_be_private(self, monkeypatch, tmp_path): + d = tmp_path / 'shared' + d.mkdir() + target = d / 'control.sock' + monkeypatch.setattr(srv, 'dev_socket_path', lambda: str(target)) + monkeypatch.setattr(os, 'geteuid', lambda: os.getuid() + 1) # "someone else's" + assert ControlServer(str(target)).start() is False + + +@needs_unix_sockets +@pytest.mark.skipif(not hasattr(socket, 'SO_PEERCRED'), reason='SO_PEERCRED is Linux-only') +class TestPermissions: + def test_peer_credentials_are_read(self, live, sock_path): + server = live() + s = _raw(sock_path) + try: + s.sendall(_req(Command.PING) + b'\n') + assert _read_line(s)['ok'] + finally: + s.close() + a, b = socket.socketpair(socket.AF_UNIX) + try: + cred = srv.peer_credentials(a) + assert cred.uid == os.geteuid() and cred.pid == os.getpid() + finally: + a.close() + b.close() + assert srv.process_groups(os.getpid()) == frozenset(os.getgroups()) + assert server.running + + @pytest.mark.skipif(hasattr(os, 'geteuid') and os.geteuid() == 0, + reason='root may always connect') + def test_a_peer_outside_the_model_is_refused(self, live, sock_path): + server = live(group=None) + # Pretend the display runs as someone else: this process is then + # neither root, the display's user, nor in its (absent) group. + server._own_uid = os.geteuid() + 12345 + s = _raw(sock_path) + try: + resp = _read_line(s) + assert resp['error']['code'] == ErrorCode.FORBIDDEN + assert s.recv(10) == b'' + finally: + s.close() + with pytest.raises(client.ControlError) as e: + client.on_demand_stop('nope', paths=[sock_path]) + assert e.value.reason == ErrorCode.FORBIDDEN + assert server.drain() == [] + + @pytest.mark.skipif(not (hasattr(os, 'geteuid') and os.geteuid() == 0), + reason='needs root to switch users (run under WSL as root, or on a Pi)') + def test_the_kernel_enforces_the_group(self, live, sock_path): + """The real deployment shape: root serves, an unprivileged user connects. + + nobody in the socket's group gets in; nobody outside it gets EACCES + from connect() -- the kernel's check, before any byte is read. + """ + import pwd + nobody = pwd.getpwnam('nobody') + allowed_gid = nobody.pw_gid + live(group=allowed_gid) + os.chmod(os.path.dirname(sock_path), 0o755) + + def try_as(gid): + r, w = os.pipe() + pid = os.fork() + if pid == 0: # child: drop to nobody with only `gid` + os.close(r) + try: + os.setgroups([]) + os.setgid(gid) + os.setuid(nobody.pw_uid) + result = json.dumps(client.ping(paths=[sock_path])) + except client.ControlError as e: + result = 'error:' + e.reason + except Exception as e: # report anything else to the parent + result = 'crash:' + repr(e) + os.write(w, result.encode()) + os._exit(0) + os.close(w) + out = b'' + while True: + chunk = os.read(r, 4096) + if not chunk: + break + out += chunk + os.close(r) + os.waitpid(pid, 0) + return out.decode() + + assert try_as(allowed_gid) == '{"pong": true}' + other_gid = allowed_gid - 1 if allowed_gid > 1 else allowed_gid + 1 + assert try_as(other_gid) == 'error:refused' + # Someone loosens the mode by hand: the kernel lets the outsider + # connect, and SO_PEERCRED still turns it away. + os.chmod(sock_path, 0o666) + assert try_as(other_gid) == 'error:forbidden' + assert try_as(allowed_gid) == '{"pong": true}' diff --git a/web_interface/blueprints/api_v3/display.py b/web_interface/blueprints/api_v3/display.py index 1ff23750..a2fc74ba 100644 --- a/web_interface/blueprints/api_v3/display.py +++ b/web_interface/blueprints/api_v3/display.py @@ -10,6 +10,7 @@ from web_interface.blueprints.api_v3 import ( ) from web_interface import display_preview import web_interface.blueprints.api_v3 as _pkg +from src.ipc import client as control_client # Read through the module rather than bound by value: tests patch these # as module attributes, and a value binding would not see the patch. # Several are also called from helpers that live in __init__, so the @@ -29,6 +30,62 @@ def _cache_manager(): return cache +#: Socket failures that only mean "this display has no socket": stopped, +#: older than the socket, Windows, or switched off. Not worth a log line. +_QUIET_SOCKET_REASONS = frozenset({'no_socket', 'disabled', 'unsupported'}) + +#: Every reason code a response may echo as ``socket_error``: the client's +#: transport reasons plus the display's ErrorCode values. Anything else is +#: reported as ``other``, so no text taken from an exception reaches a reply. +_REPORTABLE_SOCKET_REASONS = ( + 'disabled', 'unsupported', 'no_socket', 'refused', 'timeout', 'closed', + 'bad_response', 'invalid_request', + 'bad_json', 'bad_request', 'message_too_large', 'unsupported_version', + 'unknown_command', 'invalid_args', 'busy', 'forbidden', 'internal', +) + + +def _socket_reason_code(reason): + """``reason`` as one of _REPORTABLE_SOCKET_REASONS, else ``'other'``.""" + return next((code for code in _REPORTABLE_SOCKET_REASONS if code == reason), 'other') + + +def _deliver_on_demand(payload): + """Hand an on-demand request to the display: control socket, else mailbox. + + The socket (src/ipc) answers with an acknowledgement as soon as the + display has the command queued for its render thread. Any failure -- no + socket (the display is stopped or predates it), a timeout, a refusal -- + writes the file mailbox instead, exactly as before the socket existed; + the display reads it within ON_DEMAND_POLL_INTERVAL. Both carry the same + request_id, so a request that reached the display both ways (a reply + that timed out after the command was queued) is still processed once. + + Returns ``(transport, socket_error)``: ``'socket'`` and None, or + ``'mailbox'`` and the socket failure's reason code. + """ + try: + if payload['action'] == 'start': + control_client.on_demand_start( + payload['request_id'], payload.get('plugin_id'), payload.get('mode'), + payload.get('duration'), bool(payload.get('pinned', False))) + else: + control_client.on_demand_stop(payload['request_id']) + return 'socket', None + except control_client.ControlError as e: + reason = _socket_reason_code(e.reason) + if reason in _QUIET_SOCKET_REASONS: + logger.debug("On-demand %s via the mailbox: %s", payload['action'], e) + else: + logger.warning("Control socket did not take on-demand %s (%s); " + "using the mailbox", payload['action'], e) + except Exception: # never let the socket path break the route + logger.exception("Control socket client failed; using the mailbox") + reason = 'internal' + _cache_manager().set('display_on_demand_request', payload) + return 'mailbox', reason + + @api_v3.route('/display/current', methods=['GET']) def get_display_current(): """The latest display preview, as the /stream/display SSE stream sends it. @@ -192,10 +249,11 @@ def start_on_demand_display(): resolved_plugin, ) - # Post the request to the mailbox the display process polls - # (DisplayController._poll_on_demand_requests). Written before any - # service start, so a freshly started display finds it on its first poll. - cache = _cache_manager() + # Deliver the request over the control socket, or post it to the + # mailbox the display process polls (DisplayController. + # _poll_on_demand_requests). Done before any service start: a stopped + # display has no socket, so the request lands in the mailbox, where a + # freshly started display finds it on its first poll. request_id = data.get('request_id') or str(uuid.uuid4()) request_payload = { 'request_id': request_id, @@ -206,7 +264,7 @@ def start_on_demand_display(): 'pinned': pinned, 'timestamp': _pkg.time.time() } - cache.set('display_on_demand_request', request_payload) + transport, socket_error = _deliver_on_demand(request_payload) service_status = _get_display_service_status() @@ -246,8 +304,11 @@ def start_on_demand_display(): 'mode': resolved_mode, 'duration': duration, 'pinned': pinned, - 'service': service_result + 'service': service_result, + 'transport': transport, } + if socket_error: + response_data['socket_error'] = socket_error return jsonify({'status': 'success', 'data': response_data}) @api_v3.route('/display/on-demand/stop', methods=['POST']) def stop_on_demand_display(): @@ -256,29 +317,29 @@ def stop_on_demand_display(): # _coerce_to_bool: bool("false") is True, which stopped the service. stop_service = _coerce_to_bool(data.get('stop_service', False)) - # The running display reads the stop from the mailbox within - # ON_DEMAND_POLL_INTERVAL and resumes normal rotation in place - # (_clear_on_demand); nothing is restarted. - cache = _cache_manager() + # The running display takes the stop over the control socket, or reads + # it from the mailbox within ON_DEMAND_POLL_INTERVAL, and resumes normal + # rotation in place (_clear_on_demand); nothing is restarted. request_id = data.get('request_id') or str(uuid.uuid4()) request_payload = { 'request_id': request_id, 'action': 'stop', 'timestamp': _pkg.time.time() } - cache.set('display_on_demand_request', request_payload) + transport, socket_error = _deliver_on_demand(request_payload) service_result = None if stop_service: service_result = _stop_display_service() - return jsonify({ - 'status': 'success', - 'data': { - 'request_id': request_id, - 'service': service_result - } - }) + response_data = { + 'request_id': request_id, + 'service': service_result, + 'transport': transport, + } + if socket_error: + response_data['socket_error'] = socket_error + return jsonify({'status': 'success', 'data': response_data}) @api_v3.route('/display/current-status', methods=['GET']) def get_current_display_status(): """Return the display mode/plugin currently intended to be shown.